diff --git a/scripts/diffusion_bench.py b/scripts/diffusion_bench.py new file mode 100644 index 0000000000..987144a097 --- /dev/null +++ b/scripts/diffusion_bench.py @@ -0,0 +1,486 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Standalone GPU benchmark + regression harness for the Studio diffusion backend. + +Drives ``DiffusionBackend`` directly (no HTTP server) to measure load time, peak +VRAM, and generation latency for a single GGUF image model, plus an accuracy +guard: a fixed-seed image is rendered and compared (PSNR) against a stored +reference so a precision/dtype/guard regression that silently changes output is +caught, not just speed/memory. + +Two modes: + + --write-baseline PATH run once, save metrics JSON + reference.png next to it. + --compare PATH run again, diff against the baseline, exit nonzero if a + latency / VRAM / PSNR threshold is exceeded. + +torch / diffusers are imported lazily (only after argument parsing and only +inside functions) so ``--help`` works on a host without them. Not part of CPU CI; +this needs a real GPU and a downloadable model. + +Example: + python scripts/diffusion_bench.py --write-baseline outputs/diffusion_bench/baseline.json \\ + --model unsloth/Z-Image-Turbo-GGUF --gguf z-image-turbo-Q4_K_M.gguf + # ... make changes ... + python scripts/diffusion_bench.py --compare outputs/diffusion_bench/baseline.json \\ + --model unsloth/Z-Image-Turbo-GGUF --gguf z-image-turbo-Q4_K_M.gguf +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import platform +import subprocess +import sys +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Optional + +# The backend package lives at unsloth/studio/backend; this file is at +# unsloth/scripts/diffusion_bench.py. Put the backend root on sys.path so +# ``core.inference.diffusion`` imports the same way the server does. (The actual +# backend import is deferred into main() so --help never triggers torch.) +_BACKEND_ROOT = Path(__file__).resolve().parent.parent / "studio" / "backend" +if str(_BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(_BACKEND_ROOT)) + + +# ── small helpers ────────────────────────────────────────────────────────── + + +def _now_iso() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _git_commit() -> Optional[str]: + try: + out = subprocess.run( + ["git", "rev-parse", "HEAD"], + cwd = str(Path(__file__).resolve().parent), + capture_output = True, + text = True, + timeout = 10, + ) + return out.stdout.strip() or None if out.returncode == 0 else None + except Exception: + return None + + +def _percentile(values: list[float], pct: float) -> float: + """Nearest-rank percentile over a small sample (no numpy).""" + if not values: + return 0.0 + ordered = sorted(values) + rank = int(math.ceil(pct / 100.0 * len(ordered))) - 1 + rank = max(0, min(rank, len(ordered) - 1)) + return ordered[rank] + + +def _cuda_reset_peak() -> None: + import torch + if torch.cuda.is_available(): + torch.cuda.reset_peak_memory_stats() + + +def _cuda_sync() -> None: + import torch + if torch.cuda.is_available(): + torch.cuda.synchronize() + + +def _cuda_peak_alloc() -> Optional[int]: + import torch + return int(torch.cuda.max_memory_allocated()) if torch.cuda.is_available() else None + + +def _cuda_peak_reserved() -> Optional[int]: + import torch + return int(torch.cuda.max_memory_reserved()) if torch.cuda.is_available() else None + + +def _cuda_alloc() -> Optional[int]: + import torch + return int(torch.cuda.memory_allocated()) if torch.cuda.is_available() else None + + +def _gpu_name() -> Optional[str]: + try: + import torch + if torch.cuda.is_available(): + return torch.cuda.get_device_name(0) + except Exception: + pass + return None + + +def _versions() -> dict[str, Optional[str]]: + out: dict[str, Optional[str]] = {"torch": None, "diffusers": None} + try: + import torch + out["torch"] = torch.__version__ + except Exception: + pass + try: + import diffusers + out["diffusers"] = diffusers.__version__ + except Exception: + pass + return out + + +def _psnr(ref_png: Path, cand_png: Path) -> float: + """PSNR (dB) between two PNGs; inf when identical.""" + import numpy as np + from PIL import Image + + with Image.open(ref_png) as im_a: + a = np.asarray(im_a.convert("RGB"), dtype = np.float64) + with Image.open(cand_png) as im_b: + b = np.asarray(im_b.convert("RGB"), dtype = np.float64) + if a.shape != b.shape: + # Different geometry means the comparison is meaningless; report worst case. + return 0.0 + mse = float(((a - b) ** 2).mean()) + if mse == 0.0: + return math.inf + return 20.0 * math.log10(255.0) - 10.0 * math.log10(mse) + + +# ── load + generate ──────────────────────────────────────────────────────── + + +def _wait_for_load(backend: Any, timeout_s: int = 2400) -> None: + deadline = time.time() + timeout_s + last = None + while time.time() < deadline: + p = backend.load_progress() + phase = p.get("phase") + if phase != last: + last = phase + frac = p.get("fraction") or 0.0 + bt = (p.get("bytes_total") or 0) / 1e9 + print(f" load phase={phase} frac={frac:.3f} total={bt:.2f}GB", flush = True) + if phase == "ready": + return + if phase == "error": + raise RuntimeError(f"load error: {p.get('error')}") + time.sleep(2) + raise TimeoutError(f"model load did not reach ready within {timeout_s}s") + + +def _generate_once(backend: Any, args: argparse.Namespace) -> Any: + """One generation at the fixed seed; returns the first PIL image.""" + result = backend.generate( + prompt = args.prompt, + width = args.width, + height = args.height, + steps = args.steps, + guidance = args.guidance, + seed = args.seed, + batch_size = args.batch_size, + ) + images = result["images"] + return images[0] + + +def _run(args: argparse.Namespace) -> dict[str, Any]: + """Load the model, measure load + generation, render the fixed-seed image. + + Returns the metrics dict; writes the rendered image to ``args._image_out``. + """ + from core.inference.diffusion import get_diffusion_backend + + backend = get_diffusion_backend() + status: dict[str, Any] = {} + load_metrics: dict[str, Any] = {} + gen_metrics: dict[str, Any] = {} + + try: + # ── load ── + _cuda_reset_peak() + t0 = time.time() + backend.begin_load( + args.model, + gguf_filename = args.gguf, + base_repo = args.base_repo, + family_override = args.family_override, + hf_token = os.environ.get("HF_TOKEN"), + ) + _wait_for_load(backend) + _cuda_sync() + load_metrics = { + "wall_seconds": round(time.time() - t0, 2), + "peak_vram_bytes": _cuda_peak_alloc(), + "peak_reserved_bytes": _cuda_peak_reserved(), + "final_vram_bytes": _cuda_alloc(), + } + status = backend.status() + print(f" loaded: {status}", flush = True) + + # ── warmup (discarded) ── + for _ in range(max(0, args.warmup)): + _generate_once(backend, args) + + # ── measured generations (fixed seed -> deterministic) ── + _cuda_reset_peak() + latencies: list[float] = [] + first_image = None + for i in range(max(1, args.iters)): + _cuda_sync() + g0 = time.time() + image = _generate_once(backend, args) + _cuda_sync() + latencies.append(time.time() - g0) + if first_image is None: + first_image = image + print(f" gen[{i}] {latencies[-1]:.3f}s", flush = True) + + total = sum(latencies) + gen_metrics = { + "iters": len(latencies), + "warmup": max(0, args.warmup), + "latencies_s": [round(x, 4) for x in latencies], + "median_latency_s": round(_percentile(latencies, 50), 4), + "p90_latency_s": round(_percentile(latencies, 90), 4), + "images_per_sec": round((args.batch_size * len(latencies)) / total, 4) + if total > 0 + else None, + "peak_vram_bytes": _cuda_peak_alloc(), + } + + # The fixed-seed image is the accuracy anchor. + args._image_out.parent.mkdir(parents = True, exist_ok = True) + first_image.save(args._image_out) + print(f" saved image -> {args._image_out}", flush = True) + finally: + try: + backend.unload() + except Exception as exc: # noqa: BLE001 — best-effort cleanup + print(f" warn: unload failed: {exc}", flush = True) + + return { + "env": { + "timestamp": _now_iso(), + "git_commit": _git_commit(), + "python": platform.python_version(), + "platform": platform.platform(), + "versions": _versions(), + "gpu_name": _gpu_name(), + "status": status, + }, + "load": load_metrics, + "generate": gen_metrics, + "config": { + "model": args.model, + "gguf": args.gguf, + "base_repo": args.base_repo, + "family_override": args.family_override, + "prompt": args.prompt, + "width": args.width, + "height": args.height, + "steps": args.steps, + "guidance": args.guidance, + "seed": args.seed, + "batch_size": args.batch_size, + }, + } + + +# ── modes ────────────────────────────────────────────────────────────────── + + +def _write_baseline(args: argparse.Namespace) -> int: + baseline_path = Path(args.write_baseline).resolve() + ref_png = baseline_path.parent / "reference.png" + args._image_out = ref_png + + metrics = _run(args) + metrics["accuracy"] = { + "reference_png": str(ref_png), + "width": args.width, + "height": args.height, + "steps": args.steps, + "seed": args.seed, + "dtype": (metrics["env"]["status"] or {}).get("dtype"), + } + + baseline_path.parent.mkdir(parents = True, exist_ok = True) + baseline_path.write_text(json.dumps(metrics, indent = 2)) + print("\n=== BASELINE WRITTEN ===", flush = True) + print(f" json: {baseline_path}", flush = True) + print(f" reference: {ref_png}", flush = True) + print(f" load: {metrics['load']}", flush = True) + print( + f" generate: median={metrics['generate'].get('median_latency_s')}s " + f"p90={metrics['generate'].get('p90_latency_s')}s " + f"img/s={metrics['generate'].get('images_per_sec')} " + f"peak_vram={metrics['generate'].get('peak_vram_bytes')}", + flush = True, + ) + return 0 + + +def _compare(args: argparse.Namespace) -> int: + baseline_path = Path(args.compare).resolve() + baseline = json.loads(baseline_path.read_text()) + out_dir = Path(args.out_dir).resolve() + args._image_out = out_dir / "compare.png" + + # Refuse a noisy cross-hardware / cross-dtype comparison unless forced. + base_env = baseline.get("env", {}) + base_status = base_env.get("status") or {} + cur_gpu = _gpu_name() + base_gpu = base_env.get("gpu_name") + metrics = _run(args) + cur_status = metrics["env"]["status"] or {} + + mismatch = [] + if base_gpu != cur_gpu: + mismatch.append(f"gpu {base_gpu!r} -> {cur_gpu!r}") + if base_status.get("device") != cur_status.get("device"): + mismatch.append(f"device {base_status.get('device')!r} -> {cur_status.get('device')!r}") + if base_status.get("dtype") != cur_status.get("dtype"): + mismatch.append(f"dtype {base_status.get('dtype')!r} -> {cur_status.get('dtype')!r}") + if mismatch: + print("\n!! environment mismatch vs baseline: " + "; ".join(mismatch), flush = True) + if not args.force_compare: + print(" refusing noisy comparison (pass --force-compare to override).", flush = True) + return 2 + + # PSNR vs the stored reference image. The baseline stores an absolute reference_png, + # which breaks if the baseline directory was copied/moved, so fall back to reference.png + # next to the baseline JSON. A still-missing reference is a failure below, not a silent + # pass -- otherwise the benchmark would report PASS having done no image comparison. + ref_png = Path(baseline.get("accuracy", {}).get("reference_png", "")) + if not ref_png.is_file(): + ref_png = baseline_path.parent / "reference.png" + psnr = _psnr(ref_png, args._image_out) if ref_png.is_file() else float("nan") + + base_gen = baseline.get("generate", {}) + cur_gen = metrics["generate"] + base_median = base_gen.get("median_latency_s") or 0.0 + cur_median = cur_gen.get("median_latency_s") or 0.0 + latency_reg = (cur_median - base_median) / base_median if base_median > 0 else 0.0 + + base_peak = base_gen.get("peak_vram_bytes") + cur_peak = cur_gen.get("peak_vram_bytes") + vram_reg = ((cur_peak - base_peak) / base_peak) if (base_peak and cur_peak) else 0.0 + + print("\n=== REGRESSION REPORT ===", flush = True) + print(f" {'metric':<22}{'baseline':>16}{'current':>16}{'delta':>12}", flush = True) + print( + f" {'median_latency_s':<22}{base_median:>16.4f}{cur_median:>16.4f}{latency_reg * 100:>11.1f}%", + flush = True, + ) + if base_peak and cur_peak: + print( + f" {'peak_vram_MB':<22}{base_peak / 1e6:>16.1f}{cur_peak / 1e6:>16.1f}{vram_reg * 100:>11.1f}%", + flush = True, + ) + print(f" {'psnr_dB(vs ref)':<22}{'-':>16}{psnr:>16.2f}{'':>12}", flush = True) + + failures = [] + if latency_reg > args.max_latency_regression: + failures.append( + f"latency +{latency_reg * 100:.1f}% > {args.max_latency_regression * 100:.0f}%" + ) + if base_peak and cur_peak and vram_reg > args.max_vram_regression: + failures.append(f"peak VRAM +{vram_reg * 100:.1f}% > {args.max_vram_regression * 100:.0f}%") + if math.isnan(psnr): + failures.append("PSNR reference image missing; cannot verify output quality") + elif psnr < args.min_psnr: + failures.append(f"PSNR {psnr:.2f}dB < {args.min_psnr:.1f}dB (output changed)") + + if failures: + print("\n FAIL: " + "; ".join(failures), flush = True) + return 1 + print("\n PASS: no regression beyond thresholds.", flush = True) + return 0 + + +# ── cli ──────────────────────────────────────────────────────────────────── + + +def _build_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser( + description = "Benchmark + regression guard for the Studio diffusion backend.", + formatter_class = argparse.ArgumentDefaultsHelpFormatter, + ) + p.add_argument( + "--model", default = "unsloth/Z-Image-Turbo-GGUF", help = "GGUF repo id or local path" + ) + p.add_argument( + "--gguf", + default = "z-image-turbo-Q4_K_M.gguf", + help = "transformer GGUF filename inside --model", + ) + p.add_argument("--base-repo", default = None, help = "override the diffusers base repo") + p.add_argument("--family-override", default = None, help = "force a diffusion family") + p.add_argument( + "--prompt", + default = "A cozy reading nook by a rain-streaked window, warm lamplight, " + "a cat asleep on a stack of books, highly detailed", + ) + p.add_argument("--width", type = int, default = 1024) + p.add_argument("--height", type = int, default = 1024) + p.add_argument("--steps", type = int, default = 9) + p.add_argument("--guidance", type = float, default = 0.0) + p.add_argument("--seed", type = int, default = 12345, help = "fixed seed -> deterministic image") + p.add_argument("--batch-size", type = int, default = 1) + p.add_argument("--warmup", type = int, default = 1, help = "discarded warmup generations") + p.add_argument("--iters", type = int, default = 3, help = "measured generations") + p.add_argument( + "--write-baseline", + metavar = "PATH", + default = None, + help = "run once and save metrics JSON + reference.png", + ) + p.add_argument( + "--compare", metavar = "PATH", default = None, help = "run again and diff against a baseline JSON" + ) + p.add_argument( + "--max-latency-regression", + type = float, + default = 0.10, + help = "fail if median latency rises by more than this fraction", + ) + p.add_argument( + "--max-vram-regression", + type = float, + default = 0.10, + help = "fail if peak generation VRAM rises by more than this fraction", + ) + p.add_argument( + "--min-psnr", + type = float, + default = 35.0, + help = "fail if the fixed-seed image PSNR vs reference drops below this", + ) + p.add_argument( + "--force-compare", + action = "store_true", + help = "compare even when GPU/device/dtype differ from the baseline", + ) + p.add_argument( + "--out-dir", default = "outputs/diffusion_bench", help = "where compare.png is written" + ) + return p + + +def main(argv: Optional[list[str]] = None) -> int: + args = _build_parser().parse_args(argv) + if bool(args.write_baseline) == bool(args.compare): + print("error: pass exactly one of --write-baseline / --compare", file = sys.stderr) + return 2 + if args.write_baseline: + return _write_baseline(args) + return _compare(args) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 538a152c35..70f973111b 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -16,7 +16,7 @@ from __future__ import annotations import inspect import threading import time -from dataclasses import dataclass +from dataclasses import dataclass, replace from pathlib import Path from typing import Any, Optional @@ -29,6 +29,10 @@ from .diffusion_families import ( resolve_base_repo, resolve_local_gguf_child, ) +from .diffusion_device import ( + DiffusionDeviceTarget, + resolve_diffusion_device_target, +) logger = get_logger(__name__) @@ -79,15 +83,28 @@ def _estimate_eta(total_steps: int, step: int, first_step_at: float, now: float) return max(0.0, (total_steps - step) * per_step) +def _resolve_diffusion_compute_dtype(fam: Optional[DiffusionFamily], dtype: Any) -> Any: + """Promote float16 -> float32 for fp16-incompatible families (e.g. Z-Image), + whose activations overflow float16's finite range and render a black image. + Every other dtype/family passes through unchanged.""" + if fam is None or not getattr(fam, "fp16_incompatible", False): + return dtype + import torch + + return torch.float32 if dtype == torch.float16 else dtype + + class DiffusionBackend: """Holds at most one loaded diffusers pipeline. All mutations are serialised.""" def __init__(self) -> None: - # One lock serialises load / generate / unload — a generate must not run - # while the pipeline is being swapped out from under it. status() and - # load_progress() read the single state references without the lock, so - # polling never blocks a slow load. + # _lock serialises the small state mutations (the load swap, _loading, + # _load_token, _gen). status() / load_progress() / generate_progress() + # read those references WITHOUT it, so polling never blocks a slow load. self._lock = threading.Lock() + # _generate_lock serialises generations and is the ONLY lock the denoise + # holds, so a long generation never blocks status()/unload()/a new load. + self._generate_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 @@ -98,32 +115,34 @@ class DiffusionBackend: # lock, like the chat backend), so an eviction/unload can preempt a slow # load instead of blocking on the lock for the whole download. self._cancel_event = threading.Event() - # The callback mutates this and generate_progress() reads it, both without - # the lock (generate holds it for the whole call), so polling stays live. + # The cancel Event of the generation currently in flight (or None). Set + # under _lock by unload() / a superseding load to abort that specific + # denoise (its step callback flips pipe._interrupt). Per-generation rather + # than one shared flag the next generate would clear, so a cancel can't be + # lost to a racing generate nor leak onto the wrong one. + self._active_generate_cancel: Optional[threading.Event] = None + # The callback mutates _gen and generate_progress() reads it, both lock-free, + # so per-step progress polling stays live during a generation. self._gen: Optional[_GenState] = None @property def is_loaded(self) -> bool: return self._state is not None - def _pick_device_and_dtype(self) -> tuple[str, Any]: - import torch - - if torch.cuda.is_available(): - # BF16 needs Ampere+ (compute capability >= 8); pre-Ampere cards - # (Turing/Volta/Pascal) only emulate it, so use FP16 there. (Checked by - # capability, not torch.cuda.is_bf16_supported(), which returns True via - # emulation on those cards and would still pick BF16.) - dtype = torch.bfloat16 if torch.cuda.get_device_capability()[0] >= 8 else torch.float16 - return "cuda", dtype - # Intel XPU enables the Images page (CHAT_ONLY=False in hardware.py), so - # without this branch an Arc user would silently run diffusion on CPU. - if hasattr(torch, "xpu") and torch.xpu.is_available(): - return "xpu", torch.bfloat16 - mps = getattr(torch.backends, "mps", None) - if mps is not None and mps.is_available(): - return "mps", torch.float16 - return "cpu", torch.float32 + def _resolve_device_target(self, fam: Optional[DiffusionFamily]) -> DiffusionDeviceTarget: + """The device target for the current host with the family fp16 guard applied: + promote float16 -> float32 for fp16-incompatible families (Z-Image), whose + activations overflow float16 and render a black image. The capability flags + carry through unchanged; only the dtype moves.""" + target = resolve_diffusion_device_target() + effective = _resolve_diffusion_compute_dtype(fam, target.dtype) + if effective is target.dtype: + return target + logger.warning( + "diffusion.dtype_promoted: family=%s float16 -> float32 (fp16-incompatible)", + getattr(fam, "name", None), + ) + return replace(target, dtype = effective) def _resolve_gguf_path(self, repo_id: str, gguf_filename: str, hf_token: Optional[str]) -> str: local_root = Path(repo_id).expanduser() @@ -160,30 +179,62 @@ class DiffusionBackend: base, rfilename, hf_token, cancel_event = self._cancel_event ) - # ── Background load + progress ───────────────────────────────────────── + @staticmethod + def _detect_family_for_pick( + repo_id: str, gguf_filename: Optional[str], family_override: Optional[str] + ) -> Optional[DiffusionFamily]: + """Detect the family from the repo id, falling back to the combined + path/filename for a direct local .gguf pick. The frontend splits such a + pick into (parent dir, basename), so the family keyword can live only in + the filename (e.g. /models/z-image-turbo-Q4_K_M.gguf) while the parent + directory carries none; scan it too when the directory alone is + undetectable. Only used as a fallback, so remote 'org/name' picks and + explicit overrides behave exactly as before.""" + fam = detect_family(repo_id, family_override) + if fam is None and gguf_filename and not family_override: + fam = detect_family(f"{repo_id}/{gguf_filename}", family_override) + return fam - def validate_load( + def validate_load_request( self, repo_id: str, *, gguf_filename: Optional[str] = None, family_override: Optional[str] = None, ) -> DiffusionFamily: - """Cheap, pure-synchronous pre-flight: returns the resolved family or - raises ValueError. No I/O, so a caller can reject a bad request before - taking the GPU (begin_load and load_pipeline re-check on the load path).""" + """Cheap, network-free validation shared by the route (before it evicts the + chat model) and both load paths, so an unloadable pick fails BEFORE the GPU + handoff. Raises ValueError for a missing gguf_filename or undetectable + family, and ValueError/FileNotFoundError for a bad local GGUF path. Touches + no GPU, network, or state.""" if not gguf_filename: raise ValueError( "gguf_filename is required: this backend loads single-file GGUF checkpoints only." ) - fam = detect_family(repo_id, family_override) + fam = self._detect_family_for_pick(repo_id, gguf_filename, family_override) if fam is None: raise ValueError( - f"'{repo_id}' isn't a supported image-generation model. " - f"Supported: Z-Image, Qwen-Image, FLUX.1, FLUX.2-klein." + f"Could not infer a diffusion family for '{repo_id}'. Pass family_override (z-image)." ) + # Reject a bad LOCAL pick now (the same checks the load would hit later), so + # the route never evicts a working chat model for a request that can't load. + # A path-shaped repo_id (absolute / ~ / ./ / ..) is meant to be on disk, so a + # missing one is an error here; a bare "org/name" id is a remote HF repo and + # is left for the background load to resolve. + local_root = Path(repo_id).expanduser() + if local_root.exists(): + resolve_local_gguf_child(local_root, gguf_filename) + elif ( + # POSIX path-shaped, a "."/".." prefix (covers ./ ../ and their Windows .\ ..\ + # forms), a Windows separator anywhere (never present in a bare "org/name" HF + # id), or an absolute path on this OS. + repo_id.startswith(("/", "\\", "~", ".")) or "\\" in repo_id or local_root.is_absolute() + ): + raise FileNotFoundError(f"Local model path does not exist: {repo_id}") return fam + # ── Background load + progress ───────────────────────────────────────── + def begin_load( self, repo_id: str, @@ -195,7 +246,7 @@ class DiffusionBackend: cpu_offload: bool = False, ) -> dict[str, Any]: """Validate, then run the (slow) load on a daemon thread. Returns at once.""" - fam = self.validate_load( + fam = self.validate_load_request( repo_id, gguf_filename = gguf_filename, family_override = family_override ) @@ -233,7 +284,9 @@ class DiffusionBackend: # 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")) + fam = self._detect_family_for_pick( + kwargs["repo_id"], kwargs.get("gguf_filename"), kwargs.get("family_override") + ) base = _resolve_base_repo( kwargs["repo_id"], kwargs.get("base_repo"), fam, kwargs.get("hf_token") ) @@ -306,7 +359,11 @@ class DiffusionBackend: total = 0 base_files: list[str] = [] try: - if gguf_filename: + # Skip the Hub size lookup for a LOCAL gguf path: model_info(repo_id) would + # raise on a filesystem path and (caught below) skip the base-repo lookup too, + # so the companion VAE/text-encoder files would never be prefetched and would + # instead download synchronously under the load lock. + if gguf_filename and not Path(repo_id).expanduser().exists(): info = api.model_info(repo_id, files_metadata = True, token = hf_token) total += sum(s.size or 0 for s in info.siblings if s.rfilename == gguf_filename) base_info = api.model_info(base_repo, files_metadata = True, token = hf_token) @@ -347,61 +404,78 @@ class DiffusionBackend: cpu_offload: bool = False, _load_token: Optional[int] = None, ) -> dict[str, Any]: - import diffusers - - fam = self.validate_load( + # Validate first (cheap, no torch/diffusers) so a direct call with a bad + # family fails with ValueError even in a no-diffusers runtime. + fam = self.validate_load_request( repo_id, gguf_filename = gguf_filename, family_override = family_override ) base = _resolve_base_repo(repo_id, base_repo, fam, hf_token) - device, dtype = self._pick_device_and_dtype() + target = self._resolve_device_target(fam) + device, dtype = target.device, target.dtype + import diffusers + + # Signal an in-flight denoise to abort, then take _generate_lock to WAIT for + # it to actually exit before allocating the replacement: a load is about to + # claim VRAM, so unlike unload() it must not overlap a still-live pipeline. + # The cancel makes that wait ~one step (or the rest of the denoise for a + # pipeline that ignores the step callback). 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. + # Check the token BEFORE cancelling: a superseded/cancelled worker (its + # token already bumped by unload or a newer load) must not abort the + # current model's unrelated in-flight generation on its way out. if _load_token is not None and _load_token != self._load_token: raise RuntimeError("Diffusion load was cancelled.") + if self._active_generate_cancel is not None: + self._active_generate_cancel.set() + with self._generate_lock: + with self._lock: + # Re-check under the lock (we dropped it to wait on _generate_lock): an + # unload/eviction or a newer load may have superseded this one while we + # waited, and the slow VRAM-heavy build below must not run if so. + 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() + # Free the old pipeline before allocating the new one so two + # checkpoints never sit in VRAM at once. + self._unload_locked() - # Dequantise the GGUF transformer on-device; the VAE / text-encoder / - # scheduler come from the base diffusers repo (the GGUF is transformer-only). - gguf_path = self._resolve_gguf_path(repo_id, gguf_filename, hf_token) - transformer_cls = getattr(diffusers, fam.transformer_class) - transformer = transformer_cls.from_single_file( - gguf_path, - quantization_config = diffusers.GGUFQuantizationConfig(compute_dtype = dtype), - torch_dtype = dtype, - config = base, - subfolder = "transformer", - # Forward the token: the config is fetched from the (possibly gated) - # base repo before from_pretrained gets a chance to authenticate. - token = hf_token, - ) + # Dequantise the GGUF transformer on-device; the VAE / text-encoder / + # scheduler come from the base diffusers repo (GGUF is transformer-only). + gguf_path = self._resolve_gguf_path(repo_id, gguf_filename, hf_token) + transformer_cls = getattr(diffusers, fam.transformer_class) + transformer = transformer_cls.from_single_file( + gguf_path, + quantization_config = diffusers.GGUFQuantizationConfig(compute_dtype = dtype), + torch_dtype = dtype, + config = base, + subfolder = "transformer", + # Forward the token: the config is fetched from the (possibly gated) + # base repo before from_pretrained gets a chance to authenticate. + token = hf_token, + ) - pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer} - if hf_token: - pipe_kwargs["token"] = hf_token - pipeline_cls = getattr(diffusers, fam.pipeline_class) - pipe = pipeline_cls.from_pretrained(base, **pipe_kwargs) + pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer} + if hf_token: + pipe_kwargs["token"] = hf_token + pipeline_cls = getattr(diffusers, fam.pipeline_class) + pipe = pipeline_cls.from_pretrained(base, **pipe_kwargs) - use_offload = bool(cpu_offload and device == "cuda") - if use_offload: - pipe.enable_model_cpu_offload() - else: - pipe.to(device) + use_offload = bool(cpu_offload and target.supports_model_cpu_offload) + if use_offload: + pipe.enable_model_cpu_offload() + else: + pipe.to(device) - self._state = _LoadState( - pipe = pipe, - family = fam, - repo_id = repo_id, - base_repo = base, - device = device, - dtype = str(dtype).replace("torch.", ""), - cpu_offload = use_offload, - ) + self._state = _LoadState( + pipe = pipe, + family = fam, + repo_id = repo_id, + base_repo = base, + device = device, + dtype = str(dtype).replace("torch.", ""), + cpu_offload = use_offload, + ) logger.info("diffusion.loaded: repo=%s base=%s device=%s", repo_id, base, device) return self.status() @@ -422,63 +496,92 @@ class DiffusionBackend: batch_size: int = 1, ) -> dict[str, Any]: import torch - with self._lock: - state = self._state - if state is None: - raise RuntimeError("No diffusion model is loaded.") - generator = torch.Generator(device = state.device) - if seed is None: - # 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) - - kwargs: dict[str, Any] = { - "prompt": prompt, - "width": width, - "height": height, - "num_inference_steps": steps, - # Most pipelines take guidance via "guidance_scale"; Qwen-Image - # uses "true_cfg_scale" (its distilled guidance is off). - state.family.cfg_kwarg: guidance, - "generator": generator, - # Generate the whole batch in one forward pass (VRAM-heavy). All - # share this call's seed, drawn sequentially from one generator. - "num_images_per_prompt": batch_size, - } - # Pipelines vary in which kwargs they accept (a distilled pipeline may - # take neither a negative prompt nor a step callback), so only pass - # those where the signature has them. - call_params = inspect.signature(state.pipe.__call__).parameters - if negative_prompt and "negative_prompt" in call_params: - kwargs["negative_prompt"] = negative_prompt - - gen = _GenState(total_steps = steps) - - def _on_step(pipe, step_index, timestep, callback_kwargs): - now = time.time() - gen.step = step_index + 1 - if gen.first_step_at == 0.0: - gen.first_step_at = now - gen.eta_seconds = _estimate_eta(gen.total_steps, gen.step, gen.first_step_at, now) - return callback_kwargs - - if "callback_on_step_end" in call_params: - kwargs["callback_on_step_end"] = _on_step - - self._gen = gen + # A per-generation cancel Event: unload()/a superseding load set THIS event + # (registered under _lock below) to abort just this denoise. _generate_lock + # serialises generations and is the only lock the denoise holds, so a slow + # generation never blocks status()/unload()/a new load. + cancel = threading.Event() + with self._generate_lock: + with self._lock: + state = self._state + if state is None: + raise RuntimeError("No diffusion model is loaded.") + # Register under _lock so unload()/a load can signal THIS generation. + # A cancel that arrived before now either nulled _state (we raised + # above) or targets an older generation, so nothing is lost. + self._active_generate_cancel = cancel try: - images = state.pipe(**kwargs).images + # Snapshot taken: the local `state` ref keeps the pipe alive even if + # unload() nulls _state mid-denoise, so the call below needs no _lock. + generator = torch.Generator(device = state.device) + if seed is None: + # 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) + + kwargs: dict[str, Any] = { + "prompt": prompt, + "width": width, + "height": height, + "num_inference_steps": steps, + # Most pipelines take guidance via "guidance_scale"; Qwen-Image + # uses "true_cfg_scale" (its distilled guidance is off). + state.family.cfg_kwarg: guidance, + "generator": generator, + # Generate the whole batch in one forward pass (VRAM-heavy). All + # share this call's seed, drawn sequentially from one generator. + "num_images_per_prompt": batch_size, + } + # Pipelines vary in which kwargs they accept (a distilled pipeline may + # take neither a negative prompt nor a step callback), so only pass + # those where the signature has them. + call_params = inspect.signature(state.pipe.__call__).parameters + if negative_prompt and "negative_prompt" in call_params: + kwargs["negative_prompt"] = negative_prompt + + gen = _GenState(total_steps = steps) + + def _on_step(pipe, step_index, timestep, callback_kwargs): + now = time.time() + gen.step = step_index + 1 + if gen.first_step_at == 0.0: + gen.first_step_at = now + gen.eta_seconds = _estimate_eta( + gen.total_steps, gen.step, gen.first_step_at, now + ) + # Preempt a long denoise on unload/eviction or a superseding load: + # diffusers checks pipe._interrupt and stops after the current step. + if cancel.is_set(): + pipe._interrupt = True + return callback_kwargs + + if "callback_on_step_end" in call_params: + kwargs["callback_on_step_end"] = _on_step + + self._gen = gen + try: + images = state.pipe(**kwargs).images + finally: + self._gen = None + # A cancelled denoise returns early with a partial/garbage image; + # don't hand it back to be persisted. + if cancel.is_set(): + raise RuntimeError("Diffusion generation was cancelled.") + # Return the PIL images (not yet encoded): the route embeds each + # image's recipe and persists it via the gallery. + return {"images": list(images), "seed": int(seed), "repo_id": state.repo_id} finally: - self._gen = None - # Return the PIL images (not yet encoded): the route embeds each - # image's recipe and persists it via the gallery. - return {"images": list(images), "seed": int(seed), "repo_id": state.repo_id} + # Deregister so a later unload/load can't poke a finished generation + # (only if still ours — a newer generation may have replaced it). + with self._lock: + if self._active_generate_cancel is cancel: + self._active_generate_cancel = None def generate_progress(self) -> dict[str, Any]: """Live per-step progress for an in-flight generation (lock-free read).""" @@ -500,15 +603,25 @@ class DiffusionBackend: } def unload(self) -> dict[str, Any]: - # Abort an in-flight download (it runs without the lock and checks this), - # so unload/an eviction returns promptly instead of waiting it out. + # Abort an in-flight download so unload/an eviction returns promptly instead + # of waiting it out (the download runs without _lock and checks this event). self._cancel_event.set() + # Signal an in-flight denoise, then WAIT on _generate_lock for it to exit + # before freeing _state. The GPU arbiter releases diffusion ownership and + # hands the card to chat the instant this returns, so chat must not start + # allocating VRAM while the old pipeline is still resident — that transient + # overlap is an OOM risk. The cancel bounds the wait to ~one step (the same + # bound the replacement-load path in load_pipeline already relies on). with self._lock: - self._unload_locked() - # 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 + if self._active_generate_cancel is not None: + self._active_generate_cancel.set() + with self._generate_lock: + with self._lock: + self._unload_locked() + # 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_device.py b/studio/backend/core/inference/diffusion_device.py new file mode 100644 index 0000000000..8a49f1248d --- /dev/null +++ b/studio/backend/core/inference/diffusion_device.py @@ -0,0 +1,187 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Device + dtype policy for the local diffusion backend. + +torch is imported lazily inside each function so this module stays importable in +a no-torch runtime (mirrors ``diffusion.py`` / ``diffusion_families.py``). + +Studio's hardware layer reports product backends (CUDA, XPU, MLX, CPU); diffusers +runs on PyTorch devices, so Apple Silicon maps to MPS and ROCm maps to PyTorch's +``cuda`` device type. This module centralises that mapping plus the per-backend +dtype choice and the capability flags the backend keys optimisation paths off. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Optional + + +@dataclass(frozen = True) +class DiffusionDeviceTarget: + """Resolved torch device + compute dtype + per-backend capability flags.""" + + device: str + dtype: Any + backend: str + vendor: Optional[str] + supports_model_cpu_offload: bool + supports_default_torch_compile: bool + supports_pinned_transfer: bool + + +def _studio_device_is(studio_device: Any, device_type: Any, name: str) -> bool: + """True if ``studio_device`` equals ``DeviceType.`` (when that member exists).""" + member = getattr(device_type, name, None) + return member is not None and studio_device == member + + +def resolve_diffusion_device_target() -> DiffusionDeviceTarget: + """Resolve the torch device + dtype + capability flags for diffusion. + + Prefers Studio's hardware layer when importable, else probes torch directly + (CUDA -> XPU -> MPS -> CPU). On Apple Silicon Studio reports MLX/CPU when its + product backend is gated on the ``mlx`` package, but diffusers runs on + PyTorch's MPS backend, so those cases still fall through to the MPS probe. + """ + import torch + + try: + from utils.hardware import DeviceType, get_device + from utils.hardware import hardware as hardware_mod + + studio_device = get_device() + is_rocm = bool(getattr(hardware_mod, "IS_ROCM", False)) + except Exception: + DeviceType = None + studio_device = None + is_rocm = bool(getattr(getattr(torch, "version", None), "hip", None)) + + if DeviceType is not None and studio_device is not None: + if _studio_device_is(studio_device, DeviceType, "CUDA"): + if torch.cuda.is_available(): + return _cuda_or_rocm_target(torch, is_rocm = is_rocm) + return _cpu_target(torch) + if _studio_device_is(studio_device, DeviceType, "XPU"): + return _xpu_target(torch) + # MLX / CPU / anything else: diffusers uses MPS (not MLX), so fall + # through to the torch probe below, which prefers MPS over CPU. + + if torch.cuda.is_available(): + return _cuda_or_rocm_target(torch, is_rocm = is_rocm) + + xpu = getattr(torch, "xpu", None) + if xpu is not None and callable(getattr(xpu, "is_available", None)): + try: + if xpu.is_available(): + return _xpu_target(torch) + except Exception: + pass + + return _mps_or_cpu_target(torch) + + +def _cuda_or_rocm_target(torch: Any, *, is_rocm: bool) -> DiffusionDeviceTarget: + if is_rocm: + # ROCm (AMD) does not have NVIDIA's pre-Ampere bf16-emulation quirk, so + # is_bf16_supported() is trustworthy here; bf16 only when it proves it. + try: + bf16_ok = bool(torch.cuda.is_bf16_supported()) + except Exception: + bf16_ok = False + dtype = torch.bfloat16 if bf16_ok else torch.float16 + else: + # NVIDIA: bf16 needs Ampere+ (capability major >= 8). Checked by + # capability, NOT is_bf16_supported() -- pre-Ampere cards emulate bf16 + # and report it supported, but it is slow / unwanted (the #6658 fix). + try: + major = torch.cuda.get_device_capability()[0] + except Exception: + major = 0 + dtype = torch.bfloat16 if major >= 8 else torch.float16 + return DiffusionDeviceTarget( + device = "cuda", + dtype = dtype, + backend = "rocm" if is_rocm else "cuda", + vendor = "amd" if is_rocm else "nvidia", + supports_model_cpu_offload = True, + supports_default_torch_compile = not is_rocm, + supports_pinned_transfer = True, + ) + + +def _xpu_target(torch: Any) -> DiffusionDeviceTarget: + bf16_ok = False + xpu = getattr(torch, "xpu", None) + try: + bf16_ok = bool(xpu.is_bf16_supported()) if xpu is not None else False + except Exception: + bf16_ok = False + return DiffusionDeviceTarget( + device = "xpu", + dtype = torch.bfloat16 if bf16_ok else torch.float16, + backend = "xpu", + vendor = "intel", + supports_model_cpu_offload = True, + supports_default_torch_compile = False, + supports_pinned_transfer = False, + ) + + +def _mps_supports_bfloat16(torch: Any) -> bool: + """Runtime probe for usable MPS bfloat16. + + PyTorch only supports bfloat16 on MPS on macOS 14+; on older macOS a bfloat16 + op raises. Probe with a tiny compute forced to evaluate (device->host sync) + rather than guessing from the macOS / chip version. + """ + try: + x = torch.ones(2, dtype = torch.bfloat16, device = "mps") + return bool(torch.isfinite((x + x).float()).all().item()) + except Exception: + return False + + +def _mps_or_cpu_target(torch: Any) -> DiffusionDeviceTarget: + mps_available = False + try: + mps_backend = getattr(getattr(torch, "backends", None), "mps", None) + mps_available = bool( + mps_backend is not None + and callable(getattr(mps_backend, "is_available", None)) + and mps_backend.is_available() + ) + except Exception: + mps_available = False + + if mps_available: + # Prefer bfloat16; otherwise fall back to float32, NEVER silent float16. + # Modern diffusion transformers (Z-Image, FLUX.2, ...) produce activations + # far outside float16's finite range (~6.5e4) -- Z-Image's MLP + # down-projections peak near 9e5, overflowing to inf -> NaN -> a black + # image. bfloat16 (macOS 14+) shares float32's exponent range; on older + # macOS the probe fails and float32 keeps output correct (if slower). + dtype = torch.bfloat16 if _mps_supports_bfloat16(torch) else torch.float32 + return DiffusionDeviceTarget( + device = "mps", + dtype = dtype, + backend = "mps", + vendor = "apple", + supports_model_cpu_offload = False, + supports_default_torch_compile = False, + supports_pinned_transfer = False, + ) + return _cpu_target(torch) + + +def _cpu_target(torch: Any) -> DiffusionDeviceTarget: + return DiffusionDeviceTarget( + device = "cpu", + dtype = torch.float32, + backend = "cpu", + vendor = None, + supports_model_cpu_offload = False, + supports_default_torch_compile = False, + supports_pinned_transfer = False, + ) diff --git a/studio/backend/core/inference/diffusion_families.py b/studio/backend/core/inference/diffusion_families.py index 450fa6f63a..dc753f17ba 100644 --- a/studio/backend/core/inference/diffusion_families.py +++ b/studio/backend/core/inference/diffusion_families.py @@ -31,6 +31,10 @@ class DiffusionFamily: cfg_kwarg: str = "guidance_scale" # Extra lowercased substrings (besides ``name``) that map a repo id here. aliases: tuple[str, ...] = field(default_factory = tuple) + # True for families whose activations overflow float16's finite range + # (~6.5e4) and produce inf -> NaN latents -> a black image. The backend + # promotes a resolved float16 to float32 for these at load time. + fp16_incompatible: bool = False # Keyed by architecture, not per model variant: a checkpoint's specific base repo @@ -71,6 +75,8 @@ _FAMILIES: tuple[DiffusionFamily, ...] = ( transformer_class = "ZImageTransformer2DModel", base_repo = "Tongyi-MAI/Z-Image-Turbo", aliases = ("zimage", "z_image"), + # Z-Image's MLP down-projections peak near 9e5, which overflows float16. + fp16_incompatible = True, ), ) @@ -135,6 +141,6 @@ def resolve_local_gguf_child(repo_root: Path, gguf_filename: str) -> Path: child = repo_root.joinpath(*rel.parts).resolve() if child != repo_real and repo_real not in child.parents: raise ValueError("gguf_filename must resolve to a file inside the repo.") - if not child.exists(): - raise FileNotFoundError(f"'{gguf_filename}' not found under {repo_root}.") + if not child.is_file(): + raise FileNotFoundError(f"'{gguf_filename}' is not a file under {repo_root}.") return child diff --git a/studio/backend/core/inference/image_gallery.py b/studio/backend/core/inference/image_gallery.py index 66016f807a..b6b9bb2420 100644 --- a/studio/backend/core/inference/image_gallery.py +++ b/studio/backend/core/inference/image_gallery.py @@ -143,14 +143,24 @@ def list_images(limit: Optional[int] = None, offset: int = 0) -> list[dict[str, except OSError: return [] paths.sort(key = _mtime, reverse = True) - window = paths[offset:] if limit is None else paths[offset : offset + limit] + # Page over READABLE records, not raw files: filtering a foreign/corrupt PNG out of an + # already-sliced window would drop valid images that sort after it and make the route's + # has_more wrong. Read only as far as needed to fill the requested window. + # Known Phase-1 limit: this re-reads headers from the newest down to `offset+limit` on + # every page, so a deep infinite-scroll over a very large gallery (thousands of images, + # e.g. a long uncapped batch) is O(offset) header-opens per page. PIL opens are lazy + # (header only) and this runs off the event loop, so it's not a freeze; a later phase can + # switch to cursor-based paging (resume after the last-seen record) if it starts to bite. + want = None if limit is None else offset + limit records = [] - for path in window: + for path in paths: meta = _read_meta(path) if meta is None: # not one of ours (no recipe chunk) — skip continue records.append(_record(path.stem, meta)) - return records + if want is not None and len(records) >= want: + break + return records[offset:] if limit is None else records[offset : offset + limit] def delete(image_id: str) -> bool: diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 9f175a0045..e4960d5291 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -1746,6 +1746,9 @@ class GalleryImage(BaseModel): guidance: float = Field(..., description = "Guidance scale") seed: int = Field(..., description = "Seed used") batch_index: int = Field(0, description = "Position within its batch (0-based)") + batch_size: int = Field( + 1, description = "Batch size used; with batch_index it lets restore replay this image" + ) model: Optional[str] = Field(None, description = "Model repo id that produced it") created_at: float = Field(..., description = "Creation time (epoch seconds)") diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 718814ced0..aa2f23a62e 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -10211,10 +10211,11 @@ async def load_diffusion_model( backend = get_diffusion_backend() try: - # Reject an unsupported/malformed request BEFORE taking the GPU, so a - # bad image-model pick doesn't evict the user's loaded chat model only - # to 400. validate_load is pure (no I/O), so it runs inline. - backend.validate_load( + # Validate cheaply BEFORE touching the GPU: an unloadable pick (bad family, + # missing local GGUF) must not evict a working chat model and then 400. + # validate_load_request does local-path I/O, so run it off the event loop. + await asyncio.to_thread( + backend.validate_load_request, request.model_path, gguf_filename = request.gguf_filename, family_override = request.family_override, @@ -10223,7 +10224,7 @@ async def load_diffusion_model( # compete with the training subprocess for VRAM. The chat path does the # same via _guard_chat_load_against_training; this is its image sibling. _guard_diffusion_load_against_training() - # Take the GPU from the chat backend, then kick the (slow) load onto a + # Now take the GPU from the chat backend, then kick the (slow) load onto a # background thread and return at once — the client polls images/load-progress. await asyncio.to_thread(acquire_for, DIFFUSION) status_dict = await asyncio.to_thread( @@ -10236,7 +10237,7 @@ async def load_diffusion_model( cpu_offload = request.cpu_offload, ) return DiffusionStatusResponse(**status_dict) - except ValueError as exc: + except (ValueError, FileNotFoundError) as exc: raise HTTPException(status_code = 400, detail = redact_native_paths(str(exc))) except RuntimeError as exc: # A load is already in progress. @@ -10298,6 +10299,9 @@ async def generate_diffusion_image( # Position within the batch: images here share a seed + timestamp, # so the export filename needs this to stay unique. "batch_index": index, + # The batch shares one seed, so reproducing image batch_index>0 + # needs the original batch_size: persist it so restore can replay. + "batch_size": request.batch_size, "model": result.get("repo_id"), "created_at": created_at, }, diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index d5fa9f2664..c9e4bf2230 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -3081,23 +3081,45 @@ def _repo_gguf_last_modified(repo_info) -> float: # image GGUFs in its On Device list. _DIFFUSION_GGUF_ARCHS = frozenset( { - "flux", - "flux2", - "sd1", - "sd2", - "sd3", - "sdxl", - "stable_diffusion", - "lumina2", - "qwen_image", + # ONLY the families the diffusion backend can actually assemble (see + # diffusion_families._FAMILIES). Other on-device diffusion archs (SD1/2/3, + # SDXL, PixArt, Lumina2, AuraFlow, Wan, HunyuanVideo, ...) would pass this + # Images-picker filter and then fail validate_load with a 400, so they are + # deliberately excluded until the backend supports them. + "flux", # flux.1 + "flux2", # flux.2-klein + "qwen_image", # qwen-image "qwenimage", - "auraflow", - "pixart", - "hunyuan_video", - "wan", + "z_image", # z-image + "zimage", } ) +# Known diffusion / image-video GGUF archs the backend can NOT assemble yet. These +# are the GGUF general.architecture values llama.cpp also has no architecture for, +# kept in sync with core.inference.llama_cpp.LlamaCppBackend._DIFFUSION_ARCHES +# (minus the loadable set above). Tagging them with a dedicated, non-loadable task +# keeps them OUT of the chat picker -- loading one as a chat model dies with +# "unknown model architecture" -- while also keeping them out of the Images picker +# (the task is not an IMAGE_GEN_TASK), where they would 400 in validate_load. +_UNSUPPORTED_DIFFUSION_GGUF_ARCHS = frozenset( + { + "sd1", + "sd3", + "sdxl", + "aura", + "hidream", + "cosmos", + "ltxv", + "hyvid", + "wan", + "lumina2", + } +) + +# Task tag for the archs above; mirrored by the frontend NON_CHAT_TASKS gate. +_UNSUPPORTED_DIFFUSION_TASK = "image-diffusion-unsupported" + def _gguf_architecture(path: str) -> Optional[str]: """The GGUF ``general.architecture``, or None. Delegates to the shared, @@ -3111,12 +3133,21 @@ def _gguf_architecture(path: str) -> Optional[str]: def _arch_to_task(arch: Optional[str]) -> Optional[str]: if arch is None: return None - return "text-to-image" if arch.lower() in _DIFFUSION_GGUF_ARCHS else "text-generation" + a = arch.lower() + if a in _DIFFUSION_GGUF_ARCHS: + return "text-to-image" + # A diffusion arch the backend can't assemble: hide it from chat (it would die + # in llama.cpp) without surfacing it in Images (it would 400 in validate_load). + if a in _UNSUPPORTED_DIFFUSION_GGUF_ARCHS: + return _UNSUPPORTED_DIFFUSION_TASK + return "text-generation" def _repo_gguf_task(repo_info) -> Optional[str]: """HF pipeline task of a cached GGUF repo, from its architecture: - 'text-to-image' for diffusion archs, else 'text-generation' (None if unreadable).""" + 'text-to-image' for a loadable diffusion arch, the non-loadable diffusion tag + for a recognized-but-unsupported image arch, else 'text-generation' (None if + unreadable).""" try: for path in _iter_gguf_paths(Path(repo_info.repo_path)): if _is_mmproj_filename(path.name): diff --git a/studio/backend/tests/test_cached_gguf_routes.py b/studio/backend/tests/test_cached_gguf_routes.py index 6a0260c6e2..58ab41aca7 100644 --- a/studio/backend/tests/test_cached_gguf_routes.py +++ b/studio/backend/tests/test_cached_gguf_routes.py @@ -708,6 +708,31 @@ def test_gguf_download_progress_counts_quant_subdir(monkeypatch, tmp_path): assert result["progress"] == 1.0 +def test_arch_to_task_hides_unsupported_diffusion_from_chat(): + # Loadable diffusion archs -> the Images-picker task. + assert models_route._arch_to_task("flux") == "text-to-image" + assert models_route._arch_to_task("z_image") == "text-to-image" + assert models_route._arch_to_task("qwen_image") == "text-to-image" + # A real LLM arch stays a chat model; None passes through. + assert models_route._arch_to_task("llama") == "text-generation" + assert models_route._arch_to_task(None) is None + # Known-but-unsupported diffusion archs get a task that is NEITHER chat + # ("text-generation") NOR a loadable image task ("text-to-image"), so the chat + # picker hides them (they'd die in llama.cpp) and the Images picker leaves them + # out (they'd 400 in validate_load). + for arch in ("sdxl", "sd1", "sd3", "wan", "lumina2", "hidream", "cosmos"): + task = models_route._arch_to_task(arch) + assert task == models_route._UNSUPPORTED_DIFFUSION_TASK + assert task not in ("text-generation", "text-to-image") + # Drift guard: every diffusion arch llama.cpp rejects as a chat model must be + # classified here as some image task (loadable OR unsupported), never chat. + from core.inference.llama_cpp import LlamaCppBackend + + classified = models_route._DIFFUSION_GGUF_ARCHS | models_route._UNSUPPORTED_DIFFUSION_GGUF_ARCHS + missing = {a for a in LlamaCppBackend._DIFFUSION_ARCHES if a.lower() not in classified} + assert not missing, f"diffusion archs would still show in chat: {missing}" + + def test_delete_cached_refuses_diffusion_loaded_repo(monkeypatch): # The cached-delete guard refuses deleting a repo the diffusion (Images) # backend has loaded, mirroring the chat guard, so its GGUF can't be removed diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index 5fc5d23bbe..152de7ed50 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -17,7 +17,9 @@ import pytest from core.inference.diffusion import ( DiffusionBackend, + _LoadState, _base_file_downloaded, + _resolve_diffusion_compute_dtype, ) from core.inference.diffusion_families import ( detect_family, @@ -81,6 +83,11 @@ def test_resolve_local_gguf_child(tmp_path): resolve_local_gguf_child(tmp_path, "..\\secret.gguf") with pytest.raises(FileNotFoundError): resolve_local_gguf_child(tmp_path, "missing.gguf") + # A directory named like the gguf exists but isn't loadable; reject it here so + # the preflight doesn't pass and evict chat for a pick from_single_file can't load. + (tmp_path / "dir.gguf").mkdir() + with pytest.raises(FileNotFoundError): + resolve_local_gguf_child(tmp_path, "dir.gguf") def test_resolve_local_gguf_child_blocks_symlink_escape(tmp_path): @@ -206,6 +213,12 @@ def fake_runtime(monkeypatch): # The backend imports clear_gpu_cache by reference; no-op it so unload doesn't # run real hardware detection against the stubbed torch. monkeypatch.setattr("core.inference.diffusion.clear_gpu_cache", lambda: None) + # Fake the hardware layer too: resolve_diffusion_device_target() consults + # utils.hardware.get_device(), so on a real XPU/ROCm box the host would leak + # through the stubbed torch and the device tests would resolve a non-CUDA + # target. None -> fall through to the (stubbed, monkeypatchable) torch probe. + monkeypatch.setattr("utils.hardware.get_device", lambda: None, raising = False) + monkeypatch.setattr("utils.hardware.hardware.IS_ROCM", False, raising = False) _FakePipeline.last = {} _FakeTransformer.last = {} yield @@ -312,8 +325,9 @@ def test_load_without_gguf_raises(): def test_load_unknown_family_raises(): backend = DiffusionBackend() - # User-facing message: names the supported models, no internal-API jargon. - with pytest.raises(ValueError, match = "isn't a supported image-generation model"): + # load_pipeline validates via validate_load_request, which rejects an + # undetectable family before any GPU/network work. + with pytest.raises(ValueError, match = "family"): backend.load_pipeline("some/unrecognised-repo", gguf_filename = "x.gguf") @@ -447,9 +461,11 @@ def test_pick_dtype_bf16_only_on_ampere(fake_runtime, monkeypatch): backend = DiffusionBackend() monkeypatch.setattr(torch.cuda, "is_available", lambda: True, raising = False) monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: (8, 0), raising = False) - assert backend._pick_device_and_dtype() == ("cuda", torch.bfloat16) + t = backend._resolve_device_target(None) + assert (t.device, t.dtype) == ("cuda", torch.bfloat16) monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: (7, 5), raising = False) - assert backend._pick_device_and_dtype() == ("cuda", torch.float16) + t = backend._resolve_device_target(None) + assert (t.device, t.dtype) == ("cuda", torch.float16) def test_unload_sets_cancel_event(fake_runtime): @@ -504,6 +520,361 @@ def test_prefetch_downloads_gguf_and_base(monkeypatch, tmp_path): assert ("base/repo", "vae/x.safetensors") in calls +# fp16-incompatible guard + dtype promotion + + +def test_zimage_is_fp16_incompatible(): + # Only Z-Image-class families carry the guard (their activations overflow fp16). + assert detect_family("unsloth/Z-Image-Turbo-GGUF").fp16_incompatible is True + assert detect_family("unsloth/Z-Image-GGUF").fp16_incompatible is True + assert detect_family("unsloth/Qwen-Image-2512-GGUF").fp16_incompatible is False + assert detect_family("unsloth/FLUX.1-schnell-GGUF").fp16_incompatible is False + assert detect_family("unsloth/FLUX.2-klein-4B-GGUF").fp16_incompatible is False + + +def test_resolve_compute_dtype_promotes_fp16_for_zimage(fake_runtime): + torch = sys.modules["torch"] + z = detect_family("unsloth/Z-Image-GGUF") + q = detect_family("unsloth/Qwen-Image-GGUF") + # Z-Image: fp16 -> fp32; bf16 / fp32 pass through unchanged. + assert _resolve_diffusion_compute_dtype(z, torch.float16) is torch.float32 + assert _resolve_diffusion_compute_dtype(z, torch.bfloat16) is torch.bfloat16 + assert _resolve_diffusion_compute_dtype(z, torch.float32) is torch.float32 + # An fp16-compatible family (and None) keep fp16. + assert _resolve_diffusion_compute_dtype(q, torch.float16) is torch.float16 + assert _resolve_diffusion_compute_dtype(None, torch.float16) is torch.float16 + + +def test_load_promotes_fp16_to_fp32_for_zimage_only(fake_runtime, monkeypatch, tmp_path): + torch = sys.modules["torch"] + # Pre-Ampere CUDA -> the resolver picks fp16; the guard must promote Z-Image + # (and only Z-Image) to fp32 so it doesn't render a black image. + monkeypatch.setattr(torch.cuda, "is_available", lambda: True, raising = False) + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: (7, 5), raising = False) + (tmp_path / "m.gguf").write_bytes(b"x") + + z = DiffusionBackend().load_pipeline( + str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image" + ) + assert z["device"] == "cuda" and z["dtype"] == "float32" + # The promoted dtype reaches the transformer build (and thus the quant config). + assert str(_FakeTransformer.last["torch_dtype"]) == "torch.float32" + + q = DiffusionBackend().load_pipeline( + str(tmp_path), gguf_filename = "m.gguf", family_override = "qwen-image" + ) + assert q["dtype"] == "float16" # fp16-compatible family keeps fp16 on pre-Ampere + + +# Lock split + mid-denoise cancellation + + +def test_generate_lock_split_keeps_status_responsive_and_unload_waits(fake_runtime): + import threading + + backend = DiffusionBackend() + started = threading.Event() + release = threading.Event() + + class _BlockingPipe: + def __call__(self, **kwargs): + started.set() + release.wait(5) + return types.SimpleNamespace(images = [_FakeImage()]) + + fam = detect_family("unsloth/Z-Image-GGUF") + backend._state = _LoadState( + pipe = _BlockingPipe(), + family = fam, + repo_id = "r", + base_repo = "b", + device = "cpu", + dtype = "float32", + cpu_offload = False, + ) + + out: dict = {} + + def _run(): + try: + out["res"] = backend.generate(prompt = "p", steps = 4) + except Exception as exc: # noqa: BLE001 + out["exc"] = exc + + t = threading.Thread(target = _run) + t.start() + assert started.wait(5) # the denoise is in flight, holding _generate_lock + + # status() / generate_progress() read _state lock-free, so they must NOT block + # behind the denoise. + assert backend.status()["loaded"] is True + assert backend.generate_progress()["active"] is True + + # unload() signals THIS generation's cancel, then WAITS on _generate_lock for it + # to exit before freeing _state -- so when the GPU arbiter hands the card to chat + # the moment unload returns, the old pipeline is already gone (no dual-allocation + # OOM). Run it on a thread to observe that it blocks behind the live denoise. + unload_done = threading.Event() + threading.Thread(target = lambda: (backend.unload(), unload_done.set()), daemon = True).start() + assert not unload_done.wait(0.5) # blocked behind the live denoise + assert backend._state is not None # not freed while the pipeline is still live + assert backend._active_generate_cancel is not None + assert backend._active_generate_cancel.is_set() # cancel was signalled up front + + release.set() # denoise returns and releases _generate_lock + t.join(5) + assert unload_done.wait(5) # only now does unload free _state and return + assert backend.status()["loaded"] is False + # The cancelled generation raised rather than returning a now-evicted image. + assert "exc" in out and "cancelled" in str(out["exc"]).lower() + + +def test_callback_cancellation_interrupts_denoise(fake_runtime): + import threading + + backend = DiffusionBackend() + at_step0 = threading.Event() + resume = threading.Event() + + class _SteppingPipe: + def __init__(self) -> None: + self._interrupt = False + self.steps_run = 0 + + def __call__( + self, + *, + callback_on_step_end = None, + num_inference_steps = 8, + **kwargs, + ): + for i in range(num_inference_steps): + if self._interrupt: # diffusers' interrupt protocol + break + if callback_on_step_end is not None: + callback_on_step_end(self, i, 0.0, {}) + self.steps_run = i + 1 + if i == 0: + at_step0.set() + resume.wait(5) + return types.SimpleNamespace(images = [_FakeImage()]) + + pipe = _SteppingPipe() + fam = detect_family("unsloth/Z-Image-GGUF") + backend._state = _LoadState( + pipe = pipe, + family = fam, + repo_id = "r", + base_repo = "b", + device = "cpu", + dtype = "float32", + cpu_offload = False, + ) + + out: dict = {} + + def _run(): + try: + out["res"] = backend.generate(prompt = "p", steps = 8) + except Exception as exc: # noqa: BLE001 + out["exc"] = exc + + t = threading.Thread(target = _run) + t.start() + assert at_step0.wait(5) # step 0's callback ran with no cancel pending + # Simulate an eviction / superseding load signalling THIS generation's cancel. + assert backend._active_generate_cancel is not None + backend._active_generate_cancel.set() + resume.set() + t.join(5) + # The next step's callback saw the cancel, flipped pipe._interrupt, and the loop + # broke early, so the generation raised instead of returning a partial image. + assert pipe._interrupt is True + assert pipe.steps_run < 8 + assert "exc" in out and "cancelled" in str(out["exc"]).lower() + + +def test_validate_load_request(tmp_path): + backend = DiffusionBackend() + with pytest.raises(ValueError, match = "gguf_filename"): + backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF") + with pytest.raises(ValueError, match = "family"): + backend.validate_load_request("meta/Llama-3", gguf_filename = "q.gguf") + assert ( + backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "q.gguf").name + == "z-image" + ) + # A local path with a missing child fails here (before any GPU/network work). + with pytest.raises(FileNotFoundError): + backend.validate_load_request( + str(tmp_path), gguf_filename = "missing.gguf", family_override = "z-image" + ) + (tmp_path / "m.gguf").write_bytes(b"x") + assert ( + backend.validate_load_request( + str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image" + ).name + == "z-image" + ) + # A path-shaped repo_id that does not exist is rejected here (it would otherwise + # be treated as remote, evict chat, and only fail in the background load). + with pytest.raises(FileNotFoundError): + backend.validate_load_request( + "/tmp/unsloth-definitely-missing-model", + gguf_filename = "m.gguf", + family_override = "z-image", + ) + # Windows-shaped local paths (backslash separator / .\ ..\ prefixes) are also caught + # here, so a mistyped local pick on Windows can't be mistaken for a remote HF repo. + for missing in (r".\models\z-image", r"..\z-image", r"models\z-image", r"C:\models\z"): + with pytest.raises(FileNotFoundError): + backend.validate_load_request( + missing, gguf_filename = "m.gguf", family_override = "z-image" + ) + # A bare "org/name" HF id is still treated as remote (not rejected as a local path). + assert ( + backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "q.gguf").name + == "z-image" + ) + + +def test_replacement_load_waits_for_inflight_generation(fake_runtime, tmp_path): + # A superseding load must signal the in-flight generation's cancel AND wait for + # it to release _generate_lock before allocating, so two pipelines never sit in + # VRAM at once (unlike unload(), which returns promptly without waiting). + import threading + + backend = DiffusionBackend() + started = threading.Event() + release = threading.Event() + + class _BlockingPipe: + def __call__(self, **kwargs): + started.set() + release.wait(5) + return types.SimpleNamespace(images = [_FakeImage()]) + + fam = detect_family("unsloth/Z-Image-GGUF") + backend._state = _LoadState( + pipe = _BlockingPipe(), + family = fam, + repo_id = "r", + base_repo = "b", + device = "cpu", + dtype = "float32", + cpu_offload = False, + ) + + gen_out: dict = {} + + def _gen(): + try: + backend.generate(prompt = "p", steps = 4) + except Exception as exc: # noqa: BLE001 + gen_out["exc"] = exc + + gt = threading.Thread(target = _gen) + gt.start() + assert started.wait(5) # generation in flight, holding _generate_lock + + (tmp_path / "m.gguf").write_bytes(b"x") + load_done = threading.Event() + + def _load(): + backend.load_pipeline(str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image") + load_done.set() + + lt = threading.Thread(target = _load) + lt.start() + + # The load must NOT finish while the generation still holds _generate_lock; it + # has signalled the generation's cancel and is waiting to allocate. + assert not load_done.wait(0.5) + assert backend._active_generate_cancel is not None + assert backend._active_generate_cancel.is_set() + + release.set() # the blocked denoise returns; generate() sees cancel and raises + gt.join(5) + assert load_done.wait(5) # only now does the replacement allocate + assert "exc" in gen_out and "cancelled" in str(gen_out["exc"]).lower() + assert backend.status()["loaded"] is True + assert backend.status()["repo_id"] == str(tmp_path) + + +def test_superseded_load_does_not_cancel_unrelated_generation(fake_runtime, tmp_path): + # A stale worker (its _load_token already bumped by unload/a newer load) must bail on + # the token check BEFORE touching the current model's in-flight generation -- otherwise + # it would abort an unrelated denoise on its way out. + import threading + + backend = DiffusionBackend() + started = threading.Event() + release = threading.Event() + + class _BlockingPipe: + def __call__(self, **kwargs): + started.set() + release.wait(5) + return types.SimpleNamespace(images = [_FakeImage()]) + + fam = detect_family("unsloth/Z-Image-GGUF") + backend._state = _LoadState( + pipe = _BlockingPipe(), + family = fam, + repo_id = "current", + base_repo = "b", + device = "cpu", + dtype = "float32", + cpu_offload = False, + ) + backend._load_token = 7 # the "current" token + + gen_out: dict = {} + + def _gen(): + try: + backend.generate(prompt = "p", steps = 4) + except Exception as exc: # noqa: BLE001 + gen_out["exc"] = exc + + gt = threading.Thread(target = _gen) + gt.start() + assert started.wait(5) # generation in flight + cancel_event = backend._active_generate_cancel + assert cancel_event is not None + + # A superseded load arrives with a STALE token -> must raise without cancelling. + (tmp_path / "m.gguf").write_bytes(b"x") + with pytest.raises(RuntimeError, match = "cancelled"): + backend.load_pipeline( + str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", _load_token = 1 + ) + # The in-flight generation's cancel must NOT have been signalled by the stale worker. + assert not cancel_event.is_set() + + release.set() + gt.join(5) + # The generation completed normally (not cancelled). + assert "exc" not in gen_out + # The stale load did not replace the current model. + assert backend.status()["repo_id"] == "current" + + +def test_detect_family_from_local_gguf_filename(tmp_path): + # A direct local .gguf pick splits into (parent dir, basename); the family + # keyword can live only in the filename, so validate must scan it too. + (tmp_path / "z-image-turbo-Q4_K_M.gguf").write_bytes(b"w") + backend = DiffusionBackend() + fam = backend.validate_load_request(str(tmp_path), gguf_filename = "z-image-turbo-Q4_K_M.gguf") + assert fam.name == "z-image" + # An edit-checkpoint filename is still rejected even via the filename path. + (tmp_path / "FLUX.1-Kontext-dev-Q4.gguf").write_bytes(b"w") + with pytest.raises(ValueError): + backend.validate_load_request(str(tmp_path), gguf_filename = "FLUX.1-Kontext-dev-Q4.gguf") + # A bare parent dir whose name DOES carry the keyword still works (unchanged). + assert backend._detect_family_for_pick("unsloth/Z-Image-GGUF", "x.gguf", None).name == "z-image" + + def test_run_load_does_not_stamp_superseded_progress(fake_runtime, monkeypatch): # A worker whose load is superseded mid-resolve must not stamp its progress # (base_repo / expected_bytes) onto the new load's _LoadingState. diff --git a/studio/backend/tests/test_diffusion_device.py b/studio/backend/tests/test_diffusion_device.py new file mode 100644 index 0000000000..2179b16735 --- /dev/null +++ b/studio/backend/tests/test_diffusion_device.py @@ -0,0 +1,285 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Hermetic, CPU-only tests for the diffusion device/dtype resolver. + +`torch` is stubbed via a fake module so no GPU/torch is needed, and +`utils.hardware` is either stubbed (studio-layer path) or forced to fail +(torch-probe fallback path). Both paths are asserted. +""" + +from __future__ import annotations + +import sys +import types + +from core.inference import diffusion_device as dd + + +# ── Fakes ───────────────────────────────────────────────────────────── + + +class _FakeDtype: + def __init__(self, name: str) -> None: + self.name = name + + def __eq__(self, other: object) -> bool: + return isinstance(other, _FakeDtype) and other.name == self.name + + def __hash__(self) -> int: + return hash(self.name) + + def __repr__(self) -> str: # str(dtype) -> "torch.bfloat16" + return f"torch.{self.name}" + + +BF16 = _FakeDtype("bfloat16") +FP16 = _FakeDtype("float16") +FP32 = _FakeDtype("float32") + + +class _FiniteResult: + def __init__(self, finite: bool) -> None: + self._finite = finite + + def all(self) -> "_FiniteResult": + return self + + def item(self) -> bool: + return self._finite + + +class _FakeTensor: + def __init__(self, finite: bool = True) -> None: + self._finite = finite + + def __add__(self, other: object) -> "_FakeTensor": + return self + + def float(self) -> "_FakeTensor": + return self + + +def _make_torch( + *, + cuda_available: bool = False, + capability = (8, 0), + capability_raises: bool = False, + bf16_supported: bool = False, + hip = None, + mps_available: bool = False, + mps_probe: str = "pass", # "pass" | "raise" | "nonfinite" + xpu_available = None, # None -> no xpu attr; True/False -> present + xpu_bf16: bool = False, +) -> types.ModuleType: + torch = types.ModuleType("torch") + torch.bfloat16 = BF16 + torch.float16 = FP16 + torch.float32 = FP32 + torch.version = types.SimpleNamespace(hip = hip) + + def _get_cap(): + if capability_raises: + raise RuntimeError("no capability") + return capability + + torch.cuda = types.SimpleNamespace( + is_available = lambda: cuda_available, + get_device_capability = _get_cap, + is_bf16_supported = lambda: bf16_supported, + ) + + mps_ns = types.SimpleNamespace(is_available = lambda: mps_available) + torch.backends = types.SimpleNamespace(mps = mps_ns) + + def _ones(*_a, **_k): + if mps_probe == "raise": + raise RuntimeError("bf16 unsupported on this MPS") + return _FakeTensor(finite = (mps_probe == "pass")) + + torch.ones = _ones + torch.isfinite = lambda t: _FiniteResult(getattr(t, "_finite", True)) + + if xpu_available is not None: + torch.xpu = types.SimpleNamespace( + is_available = lambda: xpu_available, + is_bf16_supported = lambda: xpu_bf16, + ) + return torch + + +def _install( + monkeypatch, + torch, + *, + studio_device = None, + is_rocm = False, + hardware_fails = False, +): + """Install the fake torch and either a fake or failing `utils.hardware`.""" + monkeypatch.setitem(sys.modules, "torch", torch) + if hardware_fails: + # Force `from utils.hardware import ...` to raise -> torch-probe fallback. + monkeypatch.setitem(sys.modules, "utils.hardware", None) + return + + class _DT: + CUDA = "cuda" + XPU = "xpu" + MLX = "mlx" + CPU = "cpu" + + fake_uh = types.ModuleType("utils.hardware") + fake_uh.DeviceType = _DT + fake_uh.get_device = lambda: studio_device + fake_uh.hardware = types.SimpleNamespace(IS_ROCM = is_rocm) + monkeypatch.setitem(sys.modules, "utils.hardware", fake_uh) + + +# ── Studio-layer path ───────────────────────────────────────────────── + + +def test_cuda_ampere_bf16(monkeypatch): + torch = _make_torch(cuda_available = True, capability = (8, 0)) + _install(monkeypatch, torch, studio_device = "cuda") + t = dd.resolve_diffusion_device_target() + assert (t.device, t.dtype, t.backend, t.vendor) == ("cuda", BF16, "cuda", "nvidia") + assert ( + t.supports_model_cpu_offload + and t.supports_default_torch_compile + and t.supports_pinned_transfer + ) + + +def test_cuda_pre_ampere_fp16(monkeypatch): + torch = _make_torch(cuda_available = True, capability = (7, 5), bf16_supported = True) + _install(monkeypatch, torch, studio_device = "cuda") + t = dd.resolve_diffusion_device_target() + # is_bf16_supported() is True (emulated) but capability < 8 -> fp16. + assert t.dtype == FP16 and t.backend == "cuda" + + +def test_cuda_capability_raises_falls_back_fp16(monkeypatch): + torch = _make_torch(cuda_available = True, capability_raises = True) + _install(monkeypatch, torch, studio_device = "cuda") + t = dd.resolve_diffusion_device_target() + assert t.dtype == FP16 and t.device == "cuda" + + +def test_cuda_studio_says_cuda_but_unavailable_is_cpu(monkeypatch): + torch = _make_torch(cuda_available = False) + _install(monkeypatch, torch, studio_device = "cuda") + t = dd.resolve_diffusion_device_target() + assert t.device == "cpu" and t.dtype == FP32 + + +def test_rocm_target(monkeypatch): + torch = _make_torch(cuda_available = True, bf16_supported = True) + _install(monkeypatch, torch, studio_device = "cuda", is_rocm = True) + t = dd.resolve_diffusion_device_target() + assert (t.device, t.backend, t.vendor) == ("cuda", "rocm", "amd") + assert t.dtype == BF16 + assert t.supports_default_torch_compile is False # ROCm disables default compile + + +def test_rocm_without_bf16_uses_fp16(monkeypatch): + torch = _make_torch(cuda_available = True, bf16_supported = False) + _install(monkeypatch, torch, studio_device = "cuda", is_rocm = True) + t = dd.resolve_diffusion_device_target() + assert t.dtype == FP16 and t.backend == "rocm" + + +def test_xpu_bf16(monkeypatch): + torch = _make_torch(xpu_available = True, xpu_bf16 = True) + _install(monkeypatch, torch, studio_device = "xpu") + t = dd.resolve_diffusion_device_target() + assert (t.device, t.backend, t.vendor, t.dtype) == ("xpu", "xpu", "intel", BF16) + assert ( + t.supports_model_cpu_offload + and not t.supports_default_torch_compile + and not t.supports_pinned_transfer + ) + + +def test_xpu_without_bf16_fp16(monkeypatch): + torch = _make_torch(xpu_available = True, xpu_bf16 = False) + _install(monkeypatch, torch, studio_device = "xpu") + t = dd.resolve_diffusion_device_target() + assert t.device == "xpu" and t.dtype == FP16 + + +def test_mps_probe_pass_bf16(monkeypatch): + torch = _make_torch(mps_available = True, mps_probe = "pass") + _install(monkeypatch, torch, studio_device = "mlx") + t = dd.resolve_diffusion_device_target() + assert (t.device, t.backend, t.vendor, t.dtype) == ("mps", "mps", "apple", BF16) + assert not t.supports_model_cpu_offload + + +def test_mps_probe_raises_uses_fp32_not_fp16(monkeypatch): + torch = _make_torch(mps_available = True, mps_probe = "raise") + _install(monkeypatch, torch, studio_device = "mlx") + t = dd.resolve_diffusion_device_target() + assert t.device == "mps" and t.dtype == FP32 # strict: never silent fp16 + + +def test_mps_probe_nonfinite_uses_fp32(monkeypatch): + torch = _make_torch(mps_available = True, mps_probe = "nonfinite") + _install(monkeypatch, torch, studio_device = "mlx") + t = dd.resolve_diffusion_device_target() + assert t.device == "mps" and t.dtype == FP32 + + +def test_studio_cpu_on_apple_prefers_mps(monkeypatch): + torch = _make_torch(mps_available = True, mps_probe = "pass") + _install(monkeypatch, torch, studio_device = "cpu") # Studio reports CPU (no mlx pkg) + t = dd.resolve_diffusion_device_target() + assert t.device == "mps" and t.dtype == BF16 + + +def test_cpu_when_nothing_available(monkeypatch): + torch = _make_torch(mps_available = False) + _install(monkeypatch, torch, studio_device = "cpu") + t = dd.resolve_diffusion_device_target() + assert (t.device, t.backend, t.vendor, t.dtype) == ("cpu", "cpu", None, FP32) + assert not any( + (t.supports_model_cpu_offload, t.supports_default_torch_compile, t.supports_pinned_transfer) + ) + + +# ── torch-probe fallback path (utils.hardware import fails) ──────────── + + +def test_fallback_cuda(monkeypatch): + torch = _make_torch(cuda_available = True, capability = (9, 0)) + _install(monkeypatch, torch, hardware_fails = True) + t = dd.resolve_diffusion_device_target() + assert t.device == "cuda" and t.dtype == BF16 and t.backend == "cuda" + + +def test_fallback_rocm_via_torch_hip(monkeypatch): + torch = _make_torch(cuda_available = True, bf16_supported = True, hip = "6.2") + _install(monkeypatch, torch, hardware_fails = True) + t = dd.resolve_diffusion_device_target() + assert t.backend == "rocm" and t.vendor == "amd" + + +def test_fallback_xpu(monkeypatch): + torch = _make_torch(cuda_available = False, xpu_available = True, xpu_bf16 = True) + _install(monkeypatch, torch, hardware_fails = True) + t = dd.resolve_diffusion_device_target() + assert t.device == "xpu" and t.dtype == BF16 + + +def test_fallback_mps(monkeypatch): + torch = _make_torch(cuda_available = False, mps_available = True, mps_probe = "pass") + _install(monkeypatch, torch, hardware_fails = True) + t = dd.resolve_diffusion_device_target() + assert t.device == "mps" and t.dtype == BF16 + + +def test_fallback_cpu(monkeypatch): + torch = _make_torch(cuda_available = False, mps_available = False) + _install(monkeypatch, torch, hardware_fails = True) + t = dd.resolve_diffusion_device_target() + assert t.device == "cpu" and t.dtype == FP32 diff --git a/studio/backend/tests/test_diffusion_routes.py b/studio/backend/tests/test_diffusion_routes.py index 878de0d477..dd6d500de9 100644 --- a/studio/backend/tests/test_diffusion_routes.py +++ b/studio/backend/tests/test_diffusion_routes.py @@ -29,10 +29,23 @@ class _FakeBackend: def is_loaded(self) -> bool: return self.loaded - def validate_load(self, model_path, **kwargs): - # The route runs this cheap pre-flight before taking the GPU. The fake - # accepts anything; tests that exercise rejection patch it to raise. - return None + def validate_load_request( + self, + model_path, + *, + gguf_filename = None, + family_override = None, + ): + # Mirror the real backend's cheap validation so the route's + # validate-before-evict ordering is exercised. + from core.inference.diffusion_families import detect_family + + if not gguf_filename: + raise ValueError("gguf_filename is required.") + fam = detect_family(model_path, family_override) + if fam is None: + raise ValueError(f"Could not infer a diffusion family for '{model_path}'.") + return fam def begin_load(self, model_path, **kwargs): # The real backend loads on a thread; the fake completes instantly. @@ -251,7 +264,7 @@ def test_load_unknown_family_returns_400(client, monkeypatch): backend = _FakeBackend() # Validation runs in the pre-flight (before the GPU is taken), so that is # where an unsupported model is rejected now. - backend.validate_load = _raise + backend.validate_load_request = _raise monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend) resp = client.post( "/api/inference/images/load", json = {"model_path": "x/y", "gguf_filename": "q.gguf"} @@ -272,7 +285,7 @@ def test_load_validation_failure_does_not_evict_chat(client, monkeypatch): def _raise(*a, **k): raise ValueError("'x/y' isn't a supported image-generation model.") - backend.validate_load = _raise + backend.validate_load_request = _raise monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend) resp = client.post( "/api/inference/images/load", json = {"model_path": "x/y", "gguf_filename": "q.gguf"} @@ -325,3 +338,44 @@ def test_routes_require_auth(): app.include_router(studio_router, prefix = "/api/inference") unauth = TestClient(app) assert unauth.get("/api/inference/images/status").status_code in (401, 403) + + +def test_invalid_family_returns_400_without_evicting_chat(client): + # An undetectable family fails validation BEFORE the GPU handoff, so the + # arbiter is never acquired and a loaded chat model would not be evicted. + resp = client.post( + "/api/inference/images/load", json = {"model_path": "x/y", "gguf_filename": "q.gguf"} + ) + assert resp.status_code == 400 + assert "family" in resp.json()["detail"] + assert gpu_arbiter._owner is None + + +def test_validate_filenotfound_maps_to_400_without_eviction(client, monkeypatch): + def _raise_fnf(*a, **k): + raise FileNotFoundError("'q.gguf' not found under /models/x.") + + backend = _FakeBackend() + backend.validate_load_request = _raise_fnf + monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend) + resp = client.post( + "/api/inference/images/load", json = {"model_path": "/models/x", "gguf_filename": "q.gguf"} + ) + assert resp.status_code == 400 + assert gpu_arbiter._owner is None + + +def test_in_progress_returns_409_after_validation_passes(client, monkeypatch): + def _busy(*a, **k): + raise RuntimeError("A diffusion load is already in progress.") + + backend = _FakeBackend() + backend.begin_load = _busy + monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend) + resp = client.post( + "/api/inference/images/load", + json = {"model_path": "unsloth/Z-Image-Turbo-GGUF", "gguf_filename": "q.gguf"}, + ) + assert resp.status_code == 409 + # Validation passed first, so the GPU WAS acquired before begin_load reported busy. + assert gpu_arbiter._owner == gpu_arbiter.DIFFUSION diff --git a/studio/backend/tests/test_image_gallery.py b/studio/backend/tests/test_image_gallery.py index 883d8b8793..128a533666 100644 --- a/studio/backend/tests/test_image_gallery.py +++ b/studio/backend/tests/test_image_gallery.py @@ -119,6 +119,19 @@ def test_list_skips_foreign_pngs(tmp_path): assert [r["prompt"] for r in listed] == ["ours"] +def test_foreign_png_in_window_does_not_drop_valid_images(): + # A foreign PNG sorting INTO the requested page must not consume a window slot and + # drop a valid image that sorts after it: paging is over readable records, not files. + _save_with_mtime("p2", 100.0) + foreign = gallery.gallery_dir() / "zzz_foreign.png" + _img().save(foreign, format = "PNG") # newest by mtime (set below), sorts first + os.utime(foreign, (300.0, 300.0)) + _save_with_mtime("p1", 200.0) + # First page of 2 must still return both real images, not [p1] (foreign eating a slot). + page1 = gallery.list_images(limit = 2, offset = 0) + assert [r["prompt"] for r in page1] == ["p1", "p2"] + + 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. diff --git a/studio/frontend/src/app/routes/__root.tsx b/studio/frontend/src/app/routes/__root.tsx index e5fa6f0191..9b959c3939 100644 --- a/studio/frontend/src/app/routes/__root.tsx +++ b/studio/frontend/src/app/routes/__root.tsx @@ -32,6 +32,7 @@ import { } from "@tanstack/react-router"; import { AnimatePresence, motion } from "motion/react"; import { + lazy, Suspense, useEffect, useLayoutEffect, @@ -57,6 +58,13 @@ function RouteFallback() { ); } +// ImagesPage is mounted persistently below (not via the /images route) so an +// in-flight image batch survives leaving the tab, mirroring ChatPage. Kept lazy +// so its bundle still loads only on the first /images visit. +const ImagesPage = lazy(() => + import("@/features/images").then((m) => ({ default: m.ImagesPage })), +); + function PersonalizationSyncMount() { usePersonalizationSync(hasAuthToken()); return null; @@ -141,6 +149,17 @@ function RootLayout() { const chatSearch = isChatRoute ? liveChatSearch : frozenChatSearch; const shouldMountChat = isChatRoute || chatMounted; + // Same persistent-mount treatment for /images so a long image batch keeps + // generating when the user flips to another tab (ImagesPage reads no URL + // search, so it needs no freeze dance -- just the mount latch). Mounts lazily + // on first /images visit, then stays mounted, hidden+inert while off-route. + const isImagesRoute = pathname === "/images"; + const [imagesMounted, setImagesMounted] = useState(isImagesRoute); + if (isImagesRoute && !imagesMounted) { + setImagesMounted(true); + } + const shouldMountImages = isImagesRoute || imagesMounted; + useTrainingUnloadGuard(); // Global export driver: streams worker logs and tracks status from any route // so an export keeps running and stays visible while training / chatting. @@ -253,12 +272,30 @@ function RootLayout() { )} + {/* Same keep-alive treatment for Images so a long batch keeps + generating off-tab; `active` force-closes its body-portaled + overlays (model selector, recipe popover, aspect dropdown) so + none can bleed over another tab while hidden. */} + {shouldMountImages && ( +
+ }> + + +
+ )} {/* Use mode="popLayout" instead of "wait" to prevent UI freezes when switching from heavy pages (like Export with many checkpoints). "popLayout" allows the new route to mount immediately while the old one animates out, avoiding blocking on expensive exit renders. See issue #5850. */} - {!isChatRoute && ( + {!isChatRoute && !isImagesRoute && ( - import("@/features/images").then((m) => ({ - default: m.ImagesPage, - })), -); - +// RootLayout renders ImagesPage persistently (so an in-flight image batch is not +// cancelled when leaving the tab); this route only owns the URL + auth gate. export const Route = createRoute({ getParentRoute: () => rootRoute, path: "/images", staticData: { title: "Images" }, beforeLoad: () => requireAuth(), - component: ImagesPage, + component: () => null, }); diff --git a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx index a33f4b57f0..cc54fd6810 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx @@ -985,6 +985,16 @@ export const IMAGE_GEN_TASKS = [ "image-text-to-image", ] as const; +// Diffusion GGUF archs the Images backend can't assemble yet (SD/SDXL/PixArt/Wan/ +// ...). The backend tags them with this task so the chat picker hides them -- they +// die with "unknown model architecture" in llama.cpp -- while the Images picker, +// which filters on IMAGE_GEN_TASKS, also leaves them out (they'd 400 on load). +const UNSUPPORTED_DIFFUSION_TASK = "image-diffusion-unsupported"; + +// Tasks that must never appear as a loadable chat model: the Images-handled +// generation tasks plus the non-loadable diffusion tag above. +const NON_CHAT_TASKS: readonly string[] = [...IMAGE_GEN_TASKS, UNSUPPORTED_DIFFUSION_TASK]; + // Editing/inpaint checkpoints are tagged image-to-image but need an input image, // which the text-to-image backend rejects (mirrors its _EDIT_KEYWORDS). Hidden by // id so they don't show in the Images picker only to 400 on load. Keeping the @@ -1010,7 +1020,7 @@ function passesTaskGate( filter: HfTaskFilter, ): boolean { if (filter) return taskMatchesFilter(repoTask, filter) && !isImageEditModel(repoId); - return !(repoTask != null && (IMAGE_GEN_TASKS as readonly string[]).includes(repoTask)); + return !(repoTask != null && NON_CHAT_TASKS.includes(repoTask)); } // Module-level caches so re-mounting the popover shows results instantly diff --git a/studio/frontend/src/features/images/api.ts b/studio/frontend/src/features/images/api.ts index fb9236cb4b..949cdc0e3c 100644 --- a/studio/frontend/src/features/images/api.ts +++ b/studio/frontend/src/features/images/api.ts @@ -62,6 +62,7 @@ export interface GalleryImage { guidance: number; seed: number; batch_index: number; + batch_size: number; model: string | null; created_at: number; } diff --git a/studio/frontend/src/features/images/images-page.tsx b/studio/frontend/src/features/images/images-page.tsx index c1187bbd09..d2421ef8a1 100644 --- a/studio/frontend/src/features/images/images-page.tsx +++ b/studio/frontend/src/features/images/images-page.tsx @@ -348,12 +348,17 @@ function RecipeRow({ function RecipePopover({ image, onRestore, + active, }: { image: GalleryImage; onRestore: (image: GalleryImage) => void; + active: boolean; }) { + // Controlled + force-closed off-tab: PopoverContent portals to body, so the + // hidden/inert page wrapper can't contain it when the page is kept mounted. + const [open, setOpen] = useState(false); return ( - + setOpen(active && o)}>