diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index e7a0d5fa49..d2fcbb2058 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -90,6 +90,10 @@ class DiffusionBackend: self._lock = threading.Lock() self._state: Optional[_LoadState] = None self._loading: Optional[_LoadingState] = None + # Bumped on every begin_load and unload so a worker whose load was + # superseded (a new load) or cancelled (unload, incl. an arbiter eviction) + # neither commits its pipeline nor stamps progress onto the current load. + self._load_token = 0 # The callback mutates this and generate_progress() reads it, both without # the lock (generate holds it for the whole call), so polling stays live. self._gen: Optional[_GenState] = None @@ -136,48 +140,64 @@ class DiffusionBackend: raise ValueError( f"Could not infer a diffusion family for '{repo_id}'. Pass family_override (z-image)." ) - base = _resolve_base_repo(repo_id, base_repo, fam, hf_token) with self._lock: # Allow starting over a previously-failed load, but not over a live one. if self._loading is not None and self._loading.error is None: raise RuntimeError("A diffusion load is already in progress.") - self._loading = _LoadingState(repo_id = repo_id, base_repo = base) + self._load_token += 1 + token = self._load_token + # Seed with the family fallback; the worker resolves the real base + # (a network lookup) and updates this, so begin_load never blocks. + self._loading = _LoadingState(repo_id = repo_id, base_repo = fam.base_repo) threading.Thread( target = self._run_load, kwargs = dict( repo_id = repo_id, gguf_filename = gguf_filename, - base_repo = base, + base_repo = base_repo, family_override = family_override, hf_token = hf_token, cpu_offload = cpu_offload, + _load_token = token, ), daemon = True, ).start() return self.status() def _run_load(self, **kwargs: Any) -> None: + token = kwargs.get("_load_token") try: - # Estimate sizes on this thread (a network call) so begin_load returns - # instantly; the bar shows raw bytes until the total lands. This is the - # only writer of _loading's fields, so no lock is needed here. + # Resolve the base repo and estimate sizes on this thread (both network + # calls) so begin_load returns instantly; the bar shows raw bytes until + # the total lands. This is the only writer of _loading's fields here. + fam = detect_family(kwargs["repo_id"], kwargs.get("family_override")) + base = _resolve_base_repo(kwargs["repo_id"], kwargs.get("base_repo"), fam, kwargs.get("hf_token")) + kwargs["base_repo"] = base loading = self._loading if loading is not None: + loading.base_repo = base loading.expected_bytes = self._estimate_download_bytes( kwargs["repo_id"], kwargs.get("gguf_filename"), - kwargs["base_repo"], + base, kwargs.get("hf_token"), ) self.load_pipeline(**kwargs) with self._lock: - self._loading = None + # Only clear the marker if this load is still the current one; a + # newer begin_load (or an unload) has its own token. + if self._load_token == token: + self._loading = None except Exception as exc: # noqa: BLE001 — surfaced to the client via load_progress + # A cancelled/superseded load raised below; don't log it as a failure + # or stamp its error onto whatever load is current now. + if self._load_token != token: + return logger.error("diffusion.load_failed: %s", exc) with self._lock: - if self._loading is not None: + if self._load_token == token and self._loading is not None: self._loading.error = str(exc) def load_progress(self) -> dict[str, Any]: @@ -245,6 +265,7 @@ class DiffusionBackend: family_override: Optional[str] = None, hf_token: Optional[str] = None, cpu_offload: bool = False, + _load_token: Optional[int] = None, ) -> dict[str, Any]: import diffusers @@ -259,6 +280,12 @@ class DiffusionBackend: device, dtype = self._pick_device_and_dtype() with self._lock: + # Bail before the (slow, VRAM-heavy) build if an unload/eviction or a + # newer load superseded this one while we were resolving/downloading — + # otherwise an evicted load would resurrect a pipeline into VRAM. + if _load_token is not None and _load_token != self._load_token: + raise RuntimeError("Diffusion load was cancelled.") + # Free the old pipeline before allocating the new one so two # checkpoints never sit in VRAM at once. self._unload_locked() @@ -324,12 +351,14 @@ class DiffusionBackend: generator = torch.Generator(device = state.device) if seed is None: - # Generator.seed() seeds the generator with a fresh random value - # AND returns it, so the reported seed reproduces the image. - seed = generator.seed() + # Draw a fresh random seed but keep it within JS's safe-integer + # range (< 2**53), so the reported seed round-trips through JSON + # and actually reproduces the image (a raw 64-bit seed would lose + # precision in the browser and the recipe couldn't be replayed). + seed = generator.seed() & ((1 << 53) - 1) else: seed = int(seed) - generator.manual_seed(seed) + generator.manual_seed(seed) kwargs: dict[str, Any] = { "prompt": prompt, @@ -389,9 +418,10 @@ class DiffusionBackend: def unload(self) -> dict[str, Any]: with self._lock: self._unload_locked() - # Drop a finished/failed load marker so the next load starts clean. - if self._loading is not None and self._loading.error is not None: - self._loading = None + # Cancel any in-flight load (its worker checks this token before + # committing) and drop the marker so the next load starts clean. + self._load_token += 1 + self._loading = None return self.status() def _unload_locked(self) -> None: diff --git a/studio/backend/core/inference/diffusion_families.py b/studio/backend/core/inference/diffusion_families.py index 4ab9cdd5c6..51987061a0 100644 --- a/studio/backend/core/inference/diffusion_families.py +++ b/studio/backend/core/inference/diffusion_families.py @@ -131,8 +131,3 @@ def resolve_local_gguf_child(repo_root: Path, gguf_filename: str) -> Path: if not child.exists(): raise FileNotFoundError(f"'{gguf_filename}' not found under {repo_root}.") return child - - -def supported_families() -> list[dict]: - """Name + base repo for each known family (for status / UI listing).""" - return [{"name": fam.name, "base_repo": fam.base_repo} for fam in _FAMILIES] diff --git a/studio/backend/core/inference/gpu_arbiter.py b/studio/backend/core/inference/gpu_arbiter.py index 70335e7ea6..1b52cf775c 100644 --- a/studio/backend/core/inference/gpu_arbiter.py +++ b/studio/backend/core/inference/gpu_arbiter.py @@ -30,6 +30,8 @@ _owner: Optional[str] = None def _evict_chat() -> None: + import time + from core.inference import get_inference_backend from routes.inference import get_llama_cpp_backend @@ -42,6 +44,10 @@ def _evict_chat() -> None: # Kill the subprocess too, not just the model: its base CUDA context holds # VRAM the diffusion pipeline needs. orchestrator._shutdown_subprocess(timeout = 5.0) + # The driver reclaims the killed process's VRAM asynchronously; wait for free + # memory to settle before the diffusion pipeline allocates, mirroring the chat + # reload path — otherwise a warm chat→diffusion handoff can transiently OOM. + llama._wait_for_vram_settle(since_kill = time.monotonic()) def _evict_diffusion() -> None: diff --git a/studio/backend/core/inference/image_gallery.py b/studio/backend/core/inference/image_gallery.py index 4fb7990654..38c0ff9264 100644 --- a/studio/backend/core/inference/image_gallery.py +++ b/studio/backend/core/inference/image_gallery.py @@ -100,6 +100,12 @@ def image_b64(image_id: str) -> Optional[str]: return base64.b64encode(path.read_bytes()).decode("ascii") +# Recipe keys a gallery record must carry (the required GalleryImage fields, minus +# id/url which _record adds). A PNG missing any is treated as foreign and skipped, +# so a hand-dropped or older-schema file can't 500 the whole listing. +_REQUIRED_META = ("prompt", "width", "height", "steps", "guidance", "seed", "created_at") + + def _read_meta(path: Path) -> Optional[dict[str, Any]]: from PIL import Image @@ -114,7 +120,9 @@ def _read_meta(path: Path) -> Optional[dict[str, Any]]: meta = json.loads(raw) except (ValueError, TypeError): return None - return meta if isinstance(meta, dict) else None + if not isinstance(meta, dict) or any(k not in meta for k in _REQUIRED_META): + return None + return meta def _mtime(path: Path) -> float: diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index 7aa1a3a580..2e87c947e3 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -5,7 +5,6 @@ import asyncio import hashlib -import functools import json import os import shutil @@ -3071,55 +3070,13 @@ _DIFFUSION_GGUF_ARCHS = frozenset({ }) -# GGUF value type ids -> struct format (scalars). 8=string, 9=array handled separately. -_GGUF_SCALAR_FMT = { - 0: " Optional[str]: - """Read just general.architecture from a GGUF header by parsing the KV section - directly — orders of magnitude cheaper than GGUFReader on multi-GB files. - Cached by path (the cached file never changes).""" - import struct + """The GGUF ``general.architecture``, or None. Delegates to the shared, + bounds-checked header reader (cached by path/mtime/size).""" + from utils.models.gguf_metadata import read_gguf_general_metadata - try: - with open(path, "rb") as f: - if f.read(4) != b"GGUF": - return None - f.read(4) # version - f.read(8) # tensor count - (n_kv,) = struct.unpack(" Optional[str]: diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index bc33749d2e..df64528460 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -136,8 +136,27 @@ class _FakePipe: def enable_model_cpu_offload(self) -> None: self.offloaded = True - def __call__(self, **kwargs): - self.last_kwargs = kwargs + # Explicit signature (not just **kwargs) so generate()'s signature-gated + # guards for negative_prompt / callback_on_step_end actually take effect — + # a **kwargs-only fake would make `"negative_prompt" in signature` always False. + def __call__( + self, + *, + prompt = None, + negative_prompt = None, + callback_on_step_end = None, + guidance_scale = None, + true_cfg_scale = None, + **kwargs, + ): + self.last_kwargs = { + "prompt": prompt, + "negative_prompt": negative_prompt, + "callback_on_step_end": callback_on_step_end, + "guidance_scale": guidance_scale, + "true_cfg_scale": true_cfg_scale, + **kwargs, + } n = kwargs.get("num_images_per_prompt", 1) return types.SimpleNamespace(images = [_FakeImage() for _ in range(n)]) @@ -174,6 +193,9 @@ def fake_runtime(monkeypatch): diffusers.GGUFQuantizationConfig = lambda compute_dtype = None: ("quant", compute_dtype) diffusers.ZImagePipeline = _FakePipeline diffusers.ZImageTransformer2DModel = _FakeTransformer + # Qwen-Image too, so the true_cfg_scale cfg-kwarg path is exercisable. + diffusers.QwenImagePipeline = _FakePipeline + diffusers.QwenImageTransformer2DModel = _FakeTransformer monkeypatch.setitem(sys.modules, "torch", torch) monkeypatch.setitem(sys.modules, "diffusers", diffusers) @@ -207,10 +229,18 @@ def test_load_generate_unload_gguf(fake_runtime, tmp_path): assert _FakePipeline.last["base"] == "base/repo" assert "transformer" in _FakePipeline.last - gen = backend.generate(prompt = "a sloth", width = 512, height = 512, steps = 4, guidance = 3.0) + gen = backend.generate( + prompt = "a sloth", negative_prompt = "blurry", width = 512, height = 512, steps = 4, guidance = 3.0 + ) assert gen["seed"] == 4242 # random seed reported back assert gen["repo_id"] == str(tmp_path) # echoed so the route can record the model assert len(gen["images"]) == 1 # PIL images handed to the route for persistence + # z-image guides via guidance_scale (not true_cfg_scale); the signature-gated + # negative_prompt and per-step callback both reach the pipeline call. + call = backend._state.pipe.last_kwargs + assert call["guidance_scale"] == 3.0 and call["true_cfg_scale"] is None + assert call["negative_prompt"] == "blurry" + assert callable(call["callback_on_step_end"]) gen2 = backend.generate(prompt = "again", seed = 99) assert gen2["seed"] == 99 @@ -340,11 +370,44 @@ def test_estimate_eta(): assert _estimate_eta(8, 8, first_step_at = 100.0, now = 107.0) == 0.0 +def test_generate_qwen_uses_true_cfg_scale(fake_runtime, tmp_path): + (tmp_path / "model.gguf").write_bytes(b"weights") + backend = DiffusionBackend() + backend.load_pipeline( + str(tmp_path), gguf_filename = "model.gguf", base_repo = "Qwen/Qwen-Image", family_override = "qwen-image" + ) + backend.generate(prompt = "a sloth", guidance = 4.0) + # Qwen-Image's distilled guidance is off; the real CFG must land on true_cfg_scale. + call = backend._state.pipe.last_kwargs + assert call["true_cfg_scale"] == 4.0 and call["guidance_scale"] is None + + def test_begin_load_rejects_concurrent(monkeypatch): backend = DiffusionBackend() + # The worker resolves the base via a network lookup; stub it so the test is offline. + monkeypatch.setattr("core.inference.diffusion._hf_base_model", lambda *a, **k: None) monkeypatch.setattr(DiffusionBackend, "_estimate_download_bytes", staticmethod(lambda *a, **k: 0)) # Block the spawned worker so the load stays "in progress". monkeypatch.setattr(DiffusionBackend, "load_pipeline", lambda self, **k: __import__("time").sleep(0.2)) backend.begin_load("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "z-image-turbo-Q4_K_S.gguf") with pytest.raises(RuntimeError): backend.begin_load("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "z-image-turbo-Q4_K_S.gguf") + + +def test_unload_cancels_in_flight_load(fake_runtime): + # An unload (or an arbiter eviction, which calls unload) while a load's worker + # is still resolving/downloading must cancel it: load_pipeline sees the bumped + # token and aborts, so the evicted load never resurrects a pipeline into VRAM. + backend = DiffusionBackend() + fam = detect_family("unsloth/Z-Image-Turbo-GGUF") + token = 7 + backend._load_token = token + with pytest.raises(RuntimeError, match = "cancelled"): + # Simulate the worker reaching load_pipeline after unload bumped the token. + backend._load_token = token + 1 + backend.load_pipeline( + "unsloth/Z-Image-Turbo-GGUF", + gguf_filename = "z-image-turbo-Q4_K_S.gguf", + base_repo = fam.base_repo, + _load_token = token, + ) diff --git a/studio/backend/tests/test_image_gallery.py b/studio/backend/tests/test_image_gallery.py index b8a297c05e..883d8b8793 100644 --- a/studio/backend/tests/test_image_gallery.py +++ b/studio/backend/tests/test_image_gallery.py @@ -117,3 +117,18 @@ def test_list_skips_foreign_pngs(tmp_path): gallery.save(_img(), _meta(prompt = "ours")) listed = gallery.list_images() assert [r["prompt"] for r in listed] == ["ours"] + + +def test_list_skips_recipe_missing_required_fields(tmp_path): + # A PNG carrying our chunk but an incomplete/older-schema recipe (no seed etc.) + # must be skipped, not crash the whole listing when the route builds GalleryImage. + import json + + from PIL.PngImagePlugin import PngInfo + + info = PngInfo() + info.add_text("unsloth", json.dumps({"prompt": "partial"})) # missing width/seed/... + _img().save(gallery.gallery_dir() / "partial.png", format = "PNG", pnginfo = info) + gallery.save(_img(), _meta(prompt = "ours")) + listed = gallery.list_images() + assert [r["prompt"] for r in listed] == ["ours"] diff --git a/studio/frontend/src/features/images/images-page.tsx b/studio/frontend/src/features/images/images-page.tsx index 033762165e..d2e61431bb 100644 --- a/studio/frontend/src/features/images/images-page.tsx +++ b/studio/frontend/src/features/images/images-page.tsx @@ -80,6 +80,10 @@ const MODELS: ModelOption[] = [ // Per-model generation defaults (steps + guidance), matched by repo-id substring, // most specific first. Distilled "turbo/schnell" models want few steps and little // guidance; the full "dev" models want more steps and real CFG. +// Generation defaults when the model is unrecognised: the distilled few-step / +// no-CFG shape. Also seeds the sliders' initial state. +const DEFAULT_GEN = { steps: 9, guidance: 0 }; + const MODEL_DEFAULTS: Array<{ match: string; steps: number; guidance: number }> = [ { match: "z-image-turbo", steps: 9, guidance: 0 }, { match: "flux.1-schnell", steps: 4, guidance: 0 }, @@ -91,9 +95,9 @@ const MODEL_DEFAULTS: Array<{ match: string; steps: number; guidance: number }> function defaultsFor(repoId: string): { steps: number; guidance: number } { const id = repoId.toLowerCase(); - // Fallback (a curated entry covers every model in MODELS) is the distilled - // few-step / no-CFG shape, hit only for an unrecognised on-device image GGUF. - return MODEL_DEFAULTS.find((d) => id.includes(d.match)) ?? { steps: 9, guidance: 0 }; + // The fallback is only hit for an unrecognised on-device image GGUF; a curated + // entry covers every model in MODELS. + return MODEL_DEFAULTS.find((d) => id.includes(d.match)) ?? DEFAULT_GEN; } // Common aspect ratios (landscape; Flip gives the portrait mirror). Picking one @@ -394,8 +398,8 @@ export function ImagesPage() { const [portrait, setPortrait] = useState(false); // Z-Image-Turbo official defaults: 9 steps (= 8 DiT forwards), guidance 0 // (distilled CFG-free; a negative prompt is ignored at this guidance). - const [steps, setSteps] = useState(9); - const [guidance, setGuidance] = useState(0.0); + const [steps, setSteps] = useState(DEFAULT_GEN.steps); + const [guidance, setGuidance] = useState(DEFAULT_GEN.guidance); const [seed, setSeed] = useState(""); // Batch size = images per forward pass (VRAM-heavy); count = sequential loops. const [batchSize, setBatchSize] = useState(1); @@ -403,7 +407,9 @@ export function ImagesPage() { const [busy, setBusy] = useState(null); // {done, total} while a multi-run generation is in flight (for the button). - const [genProgress, setGenProgress] = useState<{ done: number; total: number } | null>(null); + // Number of runs finished in the current multi-run generation (null = idle). + // The total is just `count`, so it isn't stored separately. + const [genDone, setGenDone] = useState(null); // Live per-step progress (step / total + ETA) polled during generation. const [genStep, setGenStep] = useState(null); const genPollTimer = useRef | null>(null); @@ -421,6 +427,8 @@ export function ImagesPage() { const pollTimer = useRef | null>(null); // The persistent load toast's id, so each poll updates it in place (chat-style). const loadToastId = useRef(null); + // Last load-progress signature shown, so a tick that moved nothing skips the toast. + const lastLoadSig = useRef(null); const dismissLoadToast = useCallback(() => { if (loadToastId.current != null) toast.dismiss(loadToastId.current); @@ -569,12 +577,15 @@ export function ImagesPage() { useEffect(() => { void refreshStatus(); - // Stop polling if the page unmounts mid-load / mid-generate. + // Stop polling if the page unmounts mid-load / mid-generate, and dismiss the + // load toast — its poll loop is gone, so it would otherwise hang forever + // (duration: Infinity) on whatever page the user navigated to. return () => { if (pollTimer.current) clearTimeout(pollTimer.current); if (genPollTimer.current) clearInterval(genPollTimer.current); + dismissLoadToast(); }; - }, [refreshStatus]); + }, [refreshStatus, dismissLoadToast]); // Poll load-progress until the background load reaches "ready" or "error", // updating the persistent toast in place each tick. @@ -594,7 +605,13 @@ export function ImagesPage() { setBusy(null); return; } - if (loadToastId.current != null) toast(null, loadToastArgs(p, loadToastId.current)); + // Include bytes_total: the estimate lands as a 0→real jump while phase and + // bytes_downloaded hold, and the toast shows the percentage off that total. + const sig = `${p.phase}:${p.bytes_downloaded}:${p.bytes_total}`; + if (loadToastId.current != null && sig !== lastLoadSig.current) { + lastLoadSig.current = sig; + toast(null, loadToastArgs(p, loadToastId.current)); + } } catch { // Transient poll failure: keep trying. } @@ -608,6 +625,7 @@ export function ImagesPage() { setBusy("loading"); // Show the chat-style toast immediately; the poll updates it by id. dismissLoadToast(); + lastLoadSig.current = null; loadToastId.current = toast(null, loadToastArgs(IDLE_PROGRESS)); try { // Returns immediately — the load runs in the background; we poll for it. @@ -677,7 +695,7 @@ export function ImagesPage() { const h = snapDim(height); setBusy("generating"); - setGenProgress({ done: 0, total: count }); + setGenDone(0); setGenStep(null); // Poll the backend's per-step progress across the whole run (all sequential // generations), so the bar tracks the live denoising steps. @@ -712,7 +730,7 @@ export function ImagesPage() { setImages((prev) => [...res.images, ...prev]); if (res.images[0]) setSelectedId(res.images[0].id); res.images.forEach((image) => void ensureSrc(image)); - setGenProgress({ done: i + 1, total: count }); + setGenDone(i + 1); } } catch (err) { toast.error(err instanceof Error ? err.message : "Image generation failed"); @@ -720,7 +738,7 @@ export function ImagesPage() { if (genPollTimer.current) clearInterval(genPollTimer.current); genPollTimer.current = null; setBusy(null); - setGenProgress(null); + setGenDone(null); setGenStep(null); } }, [prompt, negativePrompt, width, height, steps, guidance, seed, batchSize, count, ensureSrc]); @@ -797,8 +815,8 @@ export function ImagesPage() { - - + + {busy === "generating" ? : null} - {busy === "generating" && genProgress && genProgress.total > 1 - ? `Generating ${genProgress.done}/${genProgress.total}…` + {busy === "generating" && genDone != null && count > 1 + ? `Generating ${genDone}/${count}…` : "Generate"} @@ -916,8 +934,8 @@ export function ImagesPage() {
1 - ? `Run ${genProgress.done + 1}/${genProgress.total}` + genDone != null && count > 1 + ? `Run ${genDone + 1}/${count}` : null } message="Starting…"