Revert "Made live preview and progress tracker per step based with new UI setting"

This reverts commit 7bb0b06254.
This commit is contained in:
Pablo Hellmann
2026-06-16 16:20:19 +02:00
parent 653f525ba5
commit 07d5836679
5 changed files with 60 additions and 264 deletions
-23
View File
@@ -73,18 +73,6 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict | No
shared.state.current_latent = None
shared.state.step() # increase step
shared.state.preview_job = -1 # indicate that preview image has changed
# Step-based progress and live preview push
from modules import progress
total = max(shared.state.sampling_steps, 1)
prog = round(min(1, shared.state.sampling_step / total), 2)
progress.push_progress(shared.state.sampling_step, total, prog)
if shared.opts.show_progress_every_n_steps > 0 and step % shared.opts.show_progress_every_n_steps == 0:
try:
shared.state.current_image_sampling_step = shared.state.sampling_step
shared.state.assign_current_image(image)
progress.push_live_preview(shared.state.current_image, shared.state.sampling_step, total, prog)
except Exception:
pass
debug_callback(f'Callback: step={step} timestep={timestep} image={image if image is not None else None} kwargs={list(kwargs)}')
return kwargs
@@ -241,17 +229,6 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict | No
# errors.display(e, 'Callback')
if shared.cmd_opts.profile and shared.profiler is not None:
shared.profiler.step()
# Step-based progress and live preview push
from modules import progress
total = max(shared.state.sampling_steps, 1)
prog = round(min(1, shared.state.sampling_step / total), 2)
progress.push_progress(shared.state.sampling_step, total, prog)
if shared.opts.show_progress_every_n_steps > 0 and step % shared.opts.show_progress_every_n_steps == 0:
try:
if shared.state.set_current_image() and shared.state.current_image is not None:
progress.push_live_preview(shared.state.current_image, shared.state.sampling_step, total, prog)
except Exception:
pass
t1 = time.time()
timer.process.add('callback', t1 - t0)
return kwargs
+9 -114
View File
@@ -2,9 +2,7 @@ import base64
import os
import io
import time
import asyncio
from pydantic import BaseModel, Field # pylint: disable=no-name-in-module
from starlette.websockets import WebSocket, WebSocketState
import modules.shared as shared
from modules.logger import log
@@ -18,99 +16,6 @@ debug = os.environ.get('SD_PREVIEW_DEBUG', None) is not None
debug_log = log.trace if debug else lambda *args, **kwargs: None
class PreviewManager:
def __init__(self):
self.active: list[WebSocket] = []
self._loop = None
async def connect(self, ws: WebSocket):
await ws.accept()
self._loop = asyncio.get_running_loop()
self.active.append(ws)
debug_log(f'Preview WS connect: client={ws.client.host} total={len(self.active)}')
def disconnect(self, ws: WebSocket):
try:
self.active.remove(ws)
except ValueError:
pass
debug_log(f'Preview WS disconnect: client={ws.client.host} total={len(self.active)}')
async def broadcast(self, data: dict):
disconnected = []
for ws in list(self.active):
try:
if ws.client_state == WebSocketState.CONNECTED:
await ws.send_json(data)
else:
disconnected.append(ws)
except Exception:
disconnected.append(ws)
for ws in disconnected:
try:
self.active.remove(ws)
except ValueError:
pass
def push(self, data: dict):
if self._loop is None or not self.active:
return
try:
asyncio.run_coroutine_threadsafe(self.broadcast(data), self._loop)
except Exception:
pass
preview_manager = PreviewManager()
def push_progress(step: int, steps: int, progress: float, active: bool = True):
try:
data = {
'type': 'progress',
'step': step,
'steps': steps,
'progress': progress,
'active': active,
'job': shared.state.job,
'paused': shared.state.paused,
}
preview_manager.push(data)
except Exception as e:
debug_log(f'Progress push error: {e}')
def push_live_preview(image, step: int, steps: int, progress: float):
try:
buffered = io.BytesIO()
image.save(buffered, format='jpeg', quality=60)
b64 = base64.b64encode(buffered.getvalue()).decode('ascii')
data = {
'type': 'preview',
'live_preview': f'data:image/jpeg;base64,{b64}',
'id_live_preview': shared.state.id_live_preview,
'step': step,
'steps': steps,
'progress': progress,
'job': shared.state.job,
}
preview_manager.push(data)
except Exception as e:
debug_log(f'Preview push error: {e}')
def push_complete():
try:
data = {
'type': 'complete',
'active': False,
'job': shared.state.job,
}
preview_manager.push(data)
except Exception as e:
debug_log(f'Complete push error: {e}')
def start_task(id_task):
global current_task # pylint: disable=global-statement
current_task = id_task
@@ -130,7 +35,6 @@ def finish_task(id_task):
finished_tasks.append(id_task)
if len(finished_tasks) > 16:
finished_tasks.pop(0)
push_complete()
def add_task_to_queue(id_job):
@@ -190,10 +94,15 @@ def api_progress(req: ProgressRequest):
debug_log(f'Progress: job="{shared.state.job}" active={active} progress={step}/{steps}/{progress} image={shared.state.current_image_sampling_step} request={id_live_preview} last={shared.state.id_live_preview} job={shared.state.preview_job} elapsed={elapsed:.3f}')
if active:
shared.state.set_current_image()
if shared.state.current_image is not None and req.id_live_preview != shared.state.id_live_preview:
live_preview = shared.state.live_preview_b64
if active and (req.id_live_preview != -1):
have_image = shared.state.set_current_image()
if have_image and shared.state.current_image is not None:
buffered = io.BytesIO()
shared.state.current_image.save(buffered, format='jpeg', quality=60)
b64 = base64.b64encode(buffered.getvalue())
live_preview = f'data:image/jpeg;base64,{b64.decode("ascii")}'
else:
live_preview = None
id_live_preview = shared.state.id_live_preview
@@ -220,17 +129,3 @@ def api_progress(req: ProgressRequest):
def setup_progress_api():
shared.api.add_api_route("/internal/progress", api_progress, methods=["POST"], response_model=InternalProgressResponse)
@shared.api.app.websocket("/ws/preview")
async def ws_preview(ws: WebSocket):
await preview_manager.connect(ws)
try:
while True:
data = await ws.receive_text()
if data == 'ping':
await ws.send_text('pong')
elif data == 'end':
break
except Exception:
pass
preview_manager.disconnect(ws)
+5 -14
View File
@@ -40,7 +40,6 @@ class State:
current_image = None
current_image_sampling_step = 0
id_live_preview = 0
live_preview_b64 = None
textinfo = None
prediction_type = "epsilon"
api = False
@@ -272,18 +271,13 @@ class State:
def do_set_current_image(self):
from modules import shared, images, sd_samplers_common
if self.disable_preview:
return False
# Skip if we already decoded a preview for the current sampling step
if self.current_image_sampling_step == self.sampling_step and self.current_image is not None:
if self.disable_preview or (self.preview_job == self.job_no):
return False
if (shared.opts.show_progress_type == "None") and (shared.history.last_image is not None):
self.preview_job = self.job_no
last_image = images.image_grid(shared.history.last_image)
self.assign_current_image(last_image)
self.current_image_sampling_step = self.sampling_step
self.preview_job = -1
return True
if self.current_latent is not None:
@@ -302,25 +296,22 @@ class State:
pass # ignore sigma errors
image = sd_samplers_common.samples_to_image_grid(sample)
self.assign_current_image(image)
self.preview_job = -1
return True
except Exception as e:
self.preview_job = -1
log.error(f'State image: last={self.id_live_preview} step={self.sampling_step} {e}')
display(e, 'State image')
return False
elif self.current_image is not None:
self.preview_job = self.job_no
self.current_image_sampling_step = self.sampling_step
self.assign_current_image(self.current_image)
self.preview_job = -1
return True
else:
pass
return False
def assign_current_image(self, image):
import base64
import io
self.current_image = image
self.id_live_preview += 1
buffered = io.BytesIO()
image.save(buffered, format='jpeg', quality=60)
self.live_preview_b64 = f'data:image/jpeg;base64,{base64.b64encode(buffered.getvalue()).decode("ascii")}'
+2 -2
View File
@@ -565,9 +565,9 @@ def create_settings(cmd_opts):
# --- Live Previews ---
options_templates.update(options_section(('live-preview', "Live Previews"), {
"show_progress_every_n_steps": OptionInfo(1, "Show live preview every N steps", gr.Slider, {"minimum": 1, "maximum": 20, "step": 1}, info="1 = every step, 5 = every 5th step, etc."),
"show_progress_every_n_steps": OptionInfo(1, "Live preview display period", gr.Slider, {"minimum": 0, "maximum": 20, "step": 1, "visible": False}),
"show_progress_type": OptionInfo("TAESD", "Live preview method", gr.Dropdown, {"choices": ["None", "Simple", "Approximate", "TAESD", "Full"]}),
"live_preview_refresh_period": OptionInfo(500, "Progress status update period", gr.Slider, {"minimum": 100, "maximum": 5000, "step": 25}, info="How often to poll for progress text updates in ms"),
"live_preview_refresh_period": OptionInfo(500, "Progress update period", gr.Slider, {"minimum": 0, "maximum": 5000, "step": 25}),
"taesd_variant": OptionInfo(shared_items.sd_taesd_items()[0], "TAESD variant", gr.Dropdown, {"choices": shared_items.sd_taesd_items()}),
"taesd_layers": OptionInfo(3, "TAESD decode layers", gr.Slider, {"minimum": 1, "maximum": 3, "step": 1}),
"live_preview_downscale": OptionInfo(True, "Downscale high resolution live previews"),
+44 -111
View File
@@ -13,20 +13,21 @@ export function setRefreshInterval() {
document.addEventListener('visibilitychange', () => {
if (document.hidden) refreshInterval = Math.max(2500, window.opts.live_preview_refresh_period || 1000);
else refreshInterval = window.opts.live_preview_refresh_period || 1000;
// log('refreshInterval', document.visibilityState, refreshInterval);
});
}
function pad2(x: number) {
function pad2(x) {
return x < 10 ? `0${x}` : x;
}
function formatTime(secs: number) {
function formatTime(secs) {
if (secs > 3600) return `${pad2(Math.floor(secs / 60 / 60))}:${pad2(Math.floor(secs / 60) % 60)}:${pad2(Math.floor(secs) % 60)}`;
if (secs > 60) return `${pad2(Math.floor(secs / 60))}:${pad2(Math.floor(secs) % 60)}`;
return `${Math.floor(secs)}s`;
}
export function checkPaused(state?: boolean) {
export function checkPaused(state) {
lastState.paused = state ? !state : !lastState.paused;
const t_el = document.getElementById('txt2img_pause');
const i_el = document.getElementById('img2img_pause');
@@ -83,33 +84,26 @@ export function randomId() {
return `task(${Math.random().toString(36).slice(2, 7)}${Math.random().toString(36).slice(2, 7)}${Math.random().toString(36).slice(2, 7)})`;
}
function getWebSocketUrl(): string {
const proto = location.protocol === 'https:' ? 'wss:' : 'ws:';
return `${proto}//${location.host}/ws/preview`;
}
export function requestProgress(id_task = 'undefined', progressEl: HTMLElement | null = null, galleryEl: HTMLElement | null = null, atEnd: (() => void) | null = null, onProgress: ((res: any) => void) | null = null, once = false) {
// starts sending progress requests to "/internal/progress" uri, creating progressbar above progressbarContainer element and preview inside gallery element
// Cleans up all created stuff when the task is over and calls atEnd. calls onProgress every time there is a progress update
export function requestProgress(id_task = 'undefined', progressEl = null, galleryEl = null, atEnd = null, onProgress = null, once = false) {
if (id_task) localStorage.setItem('task', id_task);
let hasStarted = false;
let dateStart = Date.now();
let prevProgress: any = null;
const parentGallery: HTMLElement | null = galleryEl ? galleryEl.parentNode as HTMLElement : null;
const parentGallery = galleryEl ? galleryEl.parentNode : null;
let livePreview: HTMLElement | undefined;
let img: HTMLImageElement;
let ws: WebSocket | null = null;
let wsReconnect: number | undefined;
let pollingTimer: number | undefined;
let id_live_preview = 0;
const initLivePreview = () => {
if (!parentGallery) return;
const footers = Array.from<any>(gradioApp().querySelectorAll('.gallery_footer'));
for (const footer of footers) {
if (footer.id !== 'gallery_footer') footer.style.display = 'none';
if (footer.id !== 'gallery_footer') footer.style.display = 'none'; // remove all footers
}
const galleries = Array.from<any>(gradioApp().querySelectorAll('.gallery_main'));
for (const gallery of galleries) {
if (gallery.id !== 'gallery_gallery') gallery.style.display = 'none';
if (gallery.id !== 'gallery_gallery') gallery.style.display = 'none'; // remove all footers
}
livePreview = document.createElement('div');
@@ -130,114 +124,50 @@ export function requestProgress(id_task = 'undefined', progressEl: HTMLElement |
debug('taskEnd:', id_task);
localStorage.removeItem('task');
setProgress();
if (ws) {
try { ws.send('end'); } catch { /* ignore */ }
try { ws.close(); } catch { /* ignore */ }
ws = null;
}
if (wsReconnect) {
clearTimeout(wsReconnect);
wsReconnect = undefined;
}
if (pollingTimer) {
clearTimeout(pollingTimer);
pollingTimer = undefined;
}
const footers = gradioApp().querySelectorAll('.gallery_footer');
for (const footer of Array.from<any>(footers)) footer.style.display = 'flex';
const galleries = gradioApp().querySelectorAll('.gallery_main');
for (const gallery of Array.from<any>(galleries)) gallery.style.display = 'flex';
if (parentGallery && livePreview && livePreview.parentNode) {
if (ok) {
const previewImg = gradioApp().querySelector('#livePreviewImage') as HTMLImageElement;
const galleryImg = gradioApp().querySelector('#control_gallery img') as HTMLImageElement;
if (previewImg?.src && galleryImg) galleryImg.src = previewImg.src;
const footers = Array.from<any>(gradioApp().querySelectorAll('.gallery_footer'));
for (const footer of footers) footer.style.display = 'flex'; // restore all footers
const galleries = Array.from<any>(gradioApp().querySelectorAll('.gallery_main'));
for (const gallery of galleries) gallery.style.display = 'flex'; // remove all galleries
try {
if (parentGallery && livePreview) {
if (ok) {
const previewImg = gradioApp().querySelector('#livePreviewImage');
const galleryImg = gradioApp().querySelector('#control_gallery img');
if (previewImg?.src && galleryImg) galleryImg.src = previewImg.src; // copy preview to gallery if everything is ok
}
parentGallery.removeChild(livePreview);
parentGallery.style.minHeight = 'unset';
parentGallery.style.maxHeight = 'unset';
parentGallery.style.overflow = 'unset';
}
parentGallery.removeChild(livePreview);
parentGallery.style.minHeight = 'unset';
parentGallery.style.maxHeight = 'unset';
parentGallery.style.overflow = 'unset';
}
} catch { /* ignore */ }
checkPaused(true);
sendNotification();
if (atEnd) atEnd();
};
const onWsMessage = (event: MessageEvent) => {
try {
const data = JSON.parse(event.data);
if (data.type === 'preview' && data.live_preview) {
if (!livePreview) initLivePreview();
if (livePreview && galleryEl && img && img.src !== data.live_preview) {
img.src = data.live_preview;
id_live_preview = data.id_live_preview;
lastState = { ...lastState, step: data.step, steps: data.steps, progress: data.progress, job: data.job };
setProgress(lastState);
if (onProgress) onProgress(lastState);
dateStart = Date.now();
prevProgress = data.progress;
}
} else if (data.type === 'progress') {
lastState = { ...lastState, step: data.step, steps: data.steps, progress: data.progress, active: data.active, paused: data.paused, job: data.job };
setProgress(lastState);
if (data.progress !== prevProgress) {
dateStart = Date.now();
prevProgress = data.progress;
}
if (onProgress) onProgress(lastState);
} else if (data.type === 'complete') {
removeLivePreview(true);
}
} catch { /* ignore */ }
};
const connectWebSocket = () => {
if (ws) return;
try {
ws = new WebSocket(getWebSocketUrl());
ws.onmessage = onWsMessage;
ws.onopen = () => {
debug('ws', 'connected');
if (wsReconnect) {
clearTimeout(wsReconnect);
wsReconnect = undefined;
}
};
ws.onclose = () => {
debug('ws', 'disconnected');
ws = null;
if (pollingTimer !== undefined) {
wsReconnect = window.setTimeout(connectWebSocket, 3000);
}
};
ws.onerror = () => {
debug('ws', 'error');
};
} catch {
wsReconnect = window.setTimeout(connectWebSocket, 3000);
}
};
const pollProgress = () => {
const startLivePreview = (taskId: string, id_live_preview: number) => {
if (window.opts.live_preview_refresh_period === 0) return;
const onProgressHandler = (res: any) => {
if (res?.debug) debug('progress:', { start: dateStart, res });
const request_id = document.hidden ? -1 : id_live_preview;
const onProgressHandler = (res) => {
if (res?.debug) debug('progress:', { start: dateStart, id: request_id, res });
lastState = res;
const elapsedFromStart = (Date.now() - dateStart) / 1000;
hasStarted = hasStarted || res.active;
if (res.completed || (!res.active && (hasStarted || once))) {
debug('progress', { end: res, reason: res.completed ? 'completed' : 'inactive' });
if (!res.paused) removeLivePreview(true);
if (!res.paused) removeLivePreview(true); // only abort if not paused
return;
}
if (elapsedFromStart > progressTimeout && !res.queued && res.progress === prevProgress) {
debug('progress', { end: res, reason: 'progressTimeout' });
if (!res.paused) removeLivePreview(false);
debug('progress', { end: res, reason: 'progressSimeout' });
if (!res.paused) removeLivePreview(false); // only abort if not paused
return;
}
if (elapsedFromStart > startTimeout && !res.queued && !res.active) {
debug('progress', { end: res, reason: 'startTimeout' });
if (!res.paused) removeLivePreview(false);
if (!res.paused) removeLivePreview(false); // only abort if not paused
return;
}
if (res.progress !== prevProgress) {
@@ -245,21 +175,24 @@ export function requestProgress(id_task = 'undefined', progressEl: HTMLElement |
prevProgress = res.progress;
}
setProgress(res);
if (res.live_preview && !livePreview) initLivePreview();
if (res.live_preview && galleryEl) {
if (img.src !== res.live_preview) img.src = res.live_preview;
id_live_preview = res.id_live_preview;
}
if (onProgress) onProgress(res);
pollingTimer = window.setTimeout(pollProgress, window.opts.live_preview_refresh_period || 500);
setTimeout(() => startLivePreview(id_task, id_live_preview), window.opts.live_preview_refresh_period || 500);
};
const onProgressErrorHandler = (err: any) => {
const onProgressErrorHandler = (err) => {
error('progress', { error: err });
removeLivePreview(false);
};
xhrPost('./internal/progress', { id_task, id_live_preview: -1 }, onProgressHandler, onProgressErrorHandler, false, 30000);
xhrPost('./internal/progress', { id_task, id_live_preview: request_id }, onProgressHandler, onProgressErrorHandler, false, 30000);
};
debug('progress', { start: dateStart });
connectWebSocket();
pollProgress();
startLivePreview(id_task, 0);
}
window.checkPaused = checkPaused;