unsloth/studio/backend/tests/test_diffusion_backend.py
Daniel Han 186f381bc6
Studio: diffusion UX polish and stronger auto policies (images + video) (#6885)
* Auto policies: deferred dense compile, video compile default, step cache and precision auto

Image dense loads with speed unset no longer sit at plain off: the load stays
bit-identical eager, and the 3rd generation in a session engages the default
compile profile plus the cuDNN attention upgrade mid-session (a one-off image
never pays the warmup, repeated use amortises it). Video dense loads resolve
straight to the default profile since a clip denoise amortises the compile
within a single run, and never to max.

Video also gains the image backend's tri-state auto policies: unset step cache
now decides from the default schedule and re-checks the actual step count per
generation, and unset precision (transformer_quant) hands the decision to the
hardware ladder instead of staying off. Memory badge reason now says plainly
that everything fits when no offload is planned.

* Rename Dtype to Precision, add the video Precision control, step cache Auto option

The images Advanced panel's Dtype row is now Precision (same control, clearer
name), and the video Advanced panel gains the matching Precision select wired
to the load route's existing transformer_quant field, gated to full-pipeline
loads the way the image control gates to GGUF. Step cache selects on both
pages gain an explicit Auto option as the default (the previous Off default
silently behaved as auto and never let anyone pin off), and the Speed and
Attention tooltips now state the deferred dense compile and the SageAttention
black-frame caveat.

* Model catalog: canonical diffusion model groups with device-aware routing

One canonical name per image/video model, its published artifacts (GGUF, FP8,
bnb-4bit, official BF16) as data, and pure routing helpers: suffix-stripped
canonical keys (owner-preserving; cross-owner merges only via explicit
aliases), group/artifact lookups, a flat back-compat options shim, load-spec
resolution replacing the pages' lookup tables, search matching over old ids
and format tokens, the GGUF fit ladder extracted from the variant expander,
and pickDefaultArtifact/pickDefaultQuant deciding what a bare group click
loads (downloaded first, then the best quality that fits 70 percent of VRAM,
GGUF as the safe fallback). Checked by npm run catalog:check, following the
i18n:check pattern.

* Picker: one canonical row per diffusion model with a format second level

The Images and Video pickers now render the curated catalog as one row per
model in Recommended: clicking loads the best artifact for the device (the
routed GGUF quant, a prequant FP8/bnb-4bit that fits, or the official BF16),
and a chevron opens the per-format list, with the GGUF row nesting the usual
quant expander. Live HF listing rows that belong to a group are deduplicated,
search collapses member repos into their group (old ids and format tokens
still match), and the On Device sections group cached member repos under the
same canonical name with the per-repo rows inside. Curated groups render from
the catalog rather than the HF listing, which finally surfaces LTX-2.3 in the
video Recommended list (its hub pipeline_tag is image-to-video, so the
text-to-video listing always missed it) and exposes the HunyuanVideo 720p
repack next to 480p.

Backend: /cached-models now tags trusted video-family repos text-to-video
instead of blanket text-to-image, and the pickers admit catalog-known
non-unsloth repos On Device, so cached Lightricks/Wan/Hunyuan pipelines
finally appear in the Video picker. Chat pickers pass no catalog and are
unchanged.

* Download formats, tab icons, plain-language train tips, 3-loop autoplay

The image Download button becomes a menu: PNG saves the original bytes with
the embedded recipe, JPEG and WebP re-encode client-side from the fetched
blob (JPEG flattened onto white). The video Download button gains MP4
(original, keeps audio), WebM and GIF; the latter two transcode server-side
from the stored MP4 via PyAV (VP9 realtime profile for WebM, ~12 fps adaptive
palette for GIF) behind a new gallery export route that 501s with a readable
message when a codec is missing.

Generated clips no longer loop forever: the player replays a clip three times
per selection, then pauses with controls up; a new generation or a refresh
gets its own three plays. The Create/Train tabs reuse the sidebar's New Chat
and Train icons (TestTubeOutlineIcon moved to a shared lib module), and every
Train tab helper text is now one plain sentence.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Keep the Create/Train tab icon and label on one line

TabsTrigger renders its children inside a plain inline span and the
Tailwind preflight gives svg display:block, so the HugeiconsIcon forced
the label onto a second line. Wrap icon plus label in their own
inline flex row inside each trigger.

* Strip -int8 and -nvfp4 prequant suffixes in the model catalog key

canonicalKeyFor already lowercases before matching, so -GGUF/-FP8 in any
case were covered; -int8 and -nvfp4 were not in the suffix table, so
such repos rendered as standalone rows in Recommended and On Device
instead of standardizing into their base-name group and routing through
pickDefaultArtifact. Added both suffixes plus case-insensitivity and
routing assertions to the catalog check.

* Standardize non-catalog picker rows to their base model name

The curated catalog already collapses its own groups, but hub listing
rows and cached repos outside the catalog (ERNIE-Image, FLUX.2-klein,
Qwen-Image-Edit-2509, FLUX.2-dev) still rendered raw ids with -GGUF /
-FP8 style suffixes in Recommended and On Device.

- model-catalog.ts: new stripArtifactSuffixesForDisplay, a
  case-preserving twin of canonicalKeyFor's stripping that keeps the
  owner prefix and original casing for display.
- pickers.tsx: recommended hub rows and the downloaded GGUF/model rows
  pass their labels through it when a catalog is present, so only the
  diffusion pickers change; chat rows keep raw ids. Click targets keep
  the full repo id, and the format badge still shows the artifact kind.
- Catalog check covers the new helper across GGUF/FP8/int8/nvfp4 in
  both cases plus no-op and suffix-only names.

* Offer official BF16/FP8 artifacts per model group and fix gallery label clipping

Model picker changes so groups are not limited to unsloth quant repos:

- model-catalog.ts: each image group that has an official vendor pipeline
  now carries its BF16 (official) artifact as the top (highest quality)
  entry - Tongyi-MAI/Z-Image-Turbo, Qwen/Qwen-Image, Qwen/Qwen-Image-2512,
  Qwen/Qwen-Image-Edit-2511, black-forest-labs/FLUX.1-dev, FLUX.1-schnell
  and FLUX.1-Kontext-dev. The LTX-2.3 video group now lists Lightricks'
  own bf16 and fp8 distilled single-file checkpoints alongside the GGUF.
  Resident sizes are set from the actual weight totals (FLUX ships a
  duplicate single-file that from_pretrained ignores, so FLUX bf16 is ~32
  GB not 54). The repos that used to be aliases are now real artifacts.
- The router already prefers the highest-quality artifact that fits the
  0.7 x GPU budget, so a datacenter GPU now defaults to official BF16
  while consumer GPUs still route to the fitting quant or GGUF. That is
  why bnb-4bit was the Z-Image-Turbo default before: it was the only
  non-GGUF artifact and it was already downloaded.
- diffusion.py: allowlist the four official image repos not previously
  trusted (qwen/qwen-image-2512, qwen/qwen-image-edit-2511,
  black-forest-labs/flux.1-schnell, flux.1-kontext-dev). All verified as
  safetensors-only diffusers model_index pipelines. The LTX-2.3
  checkpoints are already on the video trust list.
- catalog check: BF16-wins-on-datacenter, quant-wins-on-consumer, and the
  single-file load specs for the LTX-2.3 checkpoints.

Also fixes the video gallery thumbnail caption: the leading duration was
clipped by the rounded corner and selection border, so the strip now has
enough left/bottom padding to clear the curve.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* video gallery: guard export transcode against a stream-less clip

_transcode_webm and _transcode_gif indexed src.streams.video[0] before
checking the stream list, so a container with no video stream raised a bare
IndexError that the broad handlers then re-labeled as a missing libvpx or
decoder. Raise an explicit RuntimeError naming the real cause in both the
WebM and GIF paths.

* Studio: honor explicit attention/format choices, fix distilled-LTX defaults and On Device catalog routing

* Remove stray planning notes accidentally committed to the branch

* video: add transformerQuant to the load callback deps

handleLoad reads transformerQuant but omitted it from the useCallback dep array,
so after the user changes only Precision and then selects a model or clicks
Reapply, the memoized callback keeps the stale closure and loads the previous
precision. The image page's equivalent callback already lists it.

* model picker: honor the format filter when routing catalog clicks; add catalog rows to the roving list

- routedArtifactFor now scopes a group's artifacts to the active format filter
  (the same matchesFormatFilter predicate the visibility check uses) before
  pickDefaultArtifact, so a group shown only because it owns a GGUF no longer
  routes a click to a large non-GGUF download. Covers both the Recommended and
  On Device grouped paths.
- hubOptionKeys now includes the catalog-group, search-catalog-group, and grouped
  On Device row keys in exact render order, so arrow/Home/End roving reaches the
  catalog rows instead of giving them a duplicate missing id and skipping them.

* model picker: don't treat a partial base cache as downloaded

A partially-cached base repo (a cancelled download that left only some weights)
was counted as downloaded, so an On Device click routed to a fresh multi-GB
re-download instead of the complete GGUF. The picker's endpoint (/api/models/
cached-models) did not carry a partial flag at all, so a frontend-only guard
could not see it. Surface partial from that endpoint by reusing the hub inventory
scan's snapshot-partial detector, plumb it through CachedModelRepo (backend +
frontend types), and skip partial base repos when building the downloaded set.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* model picker + diffusion: drop partial/unloadable cached rows, skip defer-compile before a LoRA gen

- On Device (cached non-GGUF) rows filtered partial-download snapshots back in: sortedCachedModels
  gated on passesTaskGate + a groupForRepoId key match but, unlike downloadedSet, never checked
  c.partial, so an incomplete unsloth snapshot showed as a loadable On Device row (click errors or
  silently re-fetches multi-GB). It also admitted repos that only match the catalog by group KEY
  (a base / uncurated-quant sibling like Qwen/Qwen-Image-2512) which have no loadable artifact and
  dead-end at the trust gate. Add !c.partial and gate on artifactForRepoId (what loadSpecFor
  resolves) instead of groupForRepoId, so a cached row shows only when the backend can load it.

- Deferred speed-auto engaged the compile profile on the 3rd generation BEFORE _apply_loras. A
  compiled transformer rejects LoRA (supports_lora is False) and _apply_loras raises before its
  unchanged-selection no-op, so once compile engaged every LoRA generation on that load failed
  permanently. Skip the deferral when a LoRA is requested (compile and LoRA are mutually exclusive)
  and let it engage on a later LoRA-free generation.

* Scope the cached-model partial probe to the listed snapshot dir

list_cached_models builds each row from the largest/complete copy across HF cache
roots, but _cached_repo_partial probed is_snapshot_partial with no repo_cache_dir,
so the scan spanned every root: a stale .incomplete copy in one root would flag a
complete copy in another as partial and hide the usable model from the picker (the
click then routes to a re-download). Forward the winning snapshot's repo_path so all
three partial signals are scoped to that copy, matching the sibling inventory paths
(models/dataset cache_inventory, local_inventory).

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Do not auto-route to gated repos, prefer complete cached copies, defer compile past attached LoRA, scope group expand keys

Four fixes:
- pickDefaultArtifact's not-downloaded ladder returned the gated BF16 FLUX.1-dev / Kontext-dev
  before the open GGUF on a large GPU, so a bare group click routed to a repo the user may lack
  license/token access to. Add a gated flag and skip gated artifacts in the not-downloaded ladder
  (an already-downloaded gated artifact is still returned).
- list_cached_models picked the largest duplicate cache copy and computed partial only on it, so a
  larger partial copy shadowed a smaller complete one; since partial rows are dropped from the
  picker the usable model vanished. Prefer completeness, then size.
- the deferred-speed compile engaged on a no-LoRA generation while an adapter from a prior
  generation was still attached, baking it into the compiled graph (the later unload is swallowed
  on a compiled pipe); also defer while adapters remain attached.
- routeGroupClick's GGUF fallback toggled the context-free canonicalId while the chevron toggles
  the context-scoped expandKey, leaving the format list un-collapsible in one context, dead in the
  other, and risking cross-context expansion; thread expandKey through.

* Guard video pipeline repos from deletion, drop the always-failing LTX FP8 artifact, prefer 720p Hunyuan

Three round-6 fixes:
- cached non-GGUF video repos now surface in the Video On-Device picker with the normal delete
  action, but /delete-cached only guarded chat + the Images engine, so a loaded/loading Wan / LTX /
  Hunyuan pipeline could have its HF snapshot removed from under it. Add a VideoBackend
  loading_repo_ids accessor and a video loaded/loading guard mirroring the Images one.
- the catalog advertised Lightricks/LTX-2.3-fp8 as loadable, but the LTX-2.3 loader refuses the
  official scaled-FP8 single file (.weight_scale/.input_scale) and points to GGUF/BF16, so a pick
  routed to a ~76 GB download that always fails on load. Remove the FP8 artifact.
- pickDefaultArtifact only sorts by format, so the HunyuanVideo group's 480p (listed first) beat
  the 720p even on GPUs where 720p fits the budget. List 720p first so the fit loop prefers it and
  falls back to 480p only on smaller cards.

* diffusion: add compute int8/fp8_dynamic text-encoder quant, wire into video

Add two torchao compute text-encoder quant modes to the diffusion precision
engine, alongside the existing layerwise fp8 and weight-only nvfp4:

- int8: per-token activation + per-channel weight (torch._int_mm), with per-layer
  keep-bf16 selection. int8 degrades on large encoders unless the most
  quant-sensitive decoder blocks stay bf16, so it engages only for families with
  a measured keep-bf16 schedule (qwen-image / qwen-image-edit keep first+last 6,
  flux.2-dev keeps first 3); a family without one falls back to fp8.
- fp8_dynamic: per-row fp8 compute (torch._scaled_mm), keeping the matmul in fp8
  on the tensor cores instead of upcasting each forward like the layerwise fp8.

The selective int8 caster reuses the committed transformer-quant factory
(_make_quant_config / make_filter_fn / exclude_tokens_for_scheme) plus a small
structural first/last-N block skip, so it depends only on committed APIs.

Wire text-encoder quant into the video backend, which previously loaded the
companion encoder (Gemma3 / UMT5 / Qwen2.5-VL) dense bf16 while quantising only
the DiT. text_encoder_quant is plumbed through the load request, validation, the
load chain, the resolved record, and status, mirroring the image backend; it
applies for every load kind (the encoder is dense regardless of how the DiT was
sourced). Widen the image and video load request Literals and add the video
status field.

Tests: int8 family-schedule routing and fp8 fallback, fp8_dynamic routing,
hardware gates (int8 sm_80+, fp8_dynamic sm_89+), the structural block selection,
the real int8 filter closure (keeps the first blocks plus the vision tower /
lm_head / T5 wo dense), and the video route threading and 422 validation.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* text-encoder quant: skip the torchao modes under offload (both backends)

quantize_text_encoders applied int8-with-schedule / fp8_dynamic / nvfp4 (all torchao) to the
text encoder regardless of the offload policy. An offload placement then moves the quantized encoder
with Module.to(), which torchao tensor subclasses reject (aten._has_compatible_shallow_copy_type is
unimplemented) -- a hard crash, the same one the DiT path already skips torchao quant under offload to
avoid. Add offload_active to quantize_text_encoders and skip the torchao modes when set; layerwise fp8
is not torchao and still streams under offload. Both the video and image loaders pass
offload_active = (offload policy != none).

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* diffusion: skip non-bf16 linears for scaled_mm quant schemes

The fp8 / mxfp8 / nvfp4 schemes run on torch._scaled_mm and the fp4 / mx GEMMs,
which assert a bfloat16 input weight. On a mixed-precision DiT that keeps some
linears in fp32 for numerical stability (the Wan and Hunyuan video transformers
do this), quantize_ hits the first fp32 linear, raises, and the best-effort
wrapper swallows it to None, so the whole transformer stays dense with no error
and no speedup or memory saving.

Add a require_bf16 gate to make_filter_fn and pass it for the scaled_mm schemes
in quantize_transformer (and the fp8_dynamic text-encoder caster). The gate
skips non-bf16 linears so the scheme engages on the bf16 ones. int8 uses
torch._int_mm, which quantizes fp32/fp16 weights fine, so it leaves the gate off
and keeps its current coverage.

Verified on Wan2.2-TI2V-5B: fp8 and mxfp8 now quantize 303 linears via the
committed quantize_transformer path where they previously engaged 0.

* prequant builder: mirror the scaled-mm bf16 gate offline

The runtime DiT quantizer skips non-bf16 Linears for the scaled_mm schemes (fp8,
nvfp4, mxfp8) so the scheme engages on a mixed-precision transformer instead of
aborting on the first fp32 Linear. The offline prequant builder reused make_filter_fn
without that gate, so building an fp8/nvfp4/mxfp8 checkpoint for a mixed-precision DiT
(Wan, Hunyuan keep _keep_in_fp32_modules in fp32 even under torch_dtype=bf16) would hit
the same fp32 Linear and abort, breaking the builder's stated offline == runtime,
LPIPS-0 invariant. Thread require_bf16 = scheme in _SCALED_MM_SCHEMES through the builder,
record it in the checkpoint metadata, and verify it on load (mirrors the existing
exclude_name_tokens guard) so a future _SCALED_MM_SCHEMES change cannot silently load a
checkpoint built under the old filter.

* Keep nvfp4 fp32 linears quantised (bf16 gate is fp8/mxfp8 only)

Verified on torchao 0.17 / B200: fp8 per-row asserts 'PerRow quantization only
works for bfloat16 precision input weight' and mxfp8 asserts 'Only supporting bf16
out dtype', but NVFP4's high-precision conversion quantises an fp32 weight fine
(forward included). So the bf16 skip-gate must be fp8/mxfp8 only, not all scaled_mm
schemes -- otherwise nvfp4 leaves large fp32 projections dense, losing the intended
memory/speed gain. Rename _SCALED_MM_SCHEMES -> _REQUIRE_BF16_SCHEMES = (fp8, mxfp8)
and thread it through the runtime filter, the offline builder, and the loader
require_bf16 verification (offline == runtime preserved).

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-07-06 17:48:39 -07:00

2454 lines
103 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""CPU-only unit tests for the diffusion backend.
The family helpers are pure functions, tested directly. The backend lifecycle is
exercised with ``torch`` / ``diffusers`` stubbed via ``sys.modules`` so no real
GPU, weights, or network access is needed (sub-second, CI-friendly).
"""
from __future__ import annotations
import contextlib
import sys
import types
import pytest
from core.inference.diffusion import (
DiffusionBackend,
_LoadState,
_base_file_downloaded,
_resolve_diffusion_compute_dtype,
)
# diffusion.py imports the compile/arch patch modules LAZILY (they pull torch at module
# level, and diffusion.py must stay importable on a torchless native install). Import them
# here at collection time -- under the real torch -- so they are cached in sys.modules
# before the fake-torch fixtures swap it out; otherwise the lazy import inside load_pipeline
# would try to build them against the incomplete stub torch.
import core.inference.diffusion_eager_patches # noqa: E402,F401
import core.inference.diffusion_arch_patches # noqa: E402,F401
from core.inference.diffusion_families import (
detect_family,
resolve_base_repo,
resolve_local_gguf_child,
supported_family_names,
)
# Pure family helpers
def test_detect_family_from_repo_id():
# Detection is by architecture; Turbo/full and schnell/dev map to one family.
assert detect_family("unsloth/Z-Image-Turbo-GGUF").name == "z-image"
assert detect_family("unsloth/Z-Image-GGUF").name == "z-image"
assert detect_family("unsloth/Qwen-Image-2512-GGUF").name == "qwen-image"
assert detect_family("unsloth/FLUX.1-schnell-GGUF").name == "flux.1"
# FLUX.2-klein is its own pipeline (Qwen3 TE), distinct from FLUX.1.
klein = detect_family("unsloth/FLUX.2-klein-4B-GGUF")
assert klein.name == "flux.2-klein"
assert klein.pipeline_class == "Flux2KleinPipeline"
assert klein.cfg_kwarg == "guidance_scale"
# Both klein sizes share the one family (base repo resolved per-variant).
assert detect_family("unsloth/FLUX.2-klein-9B-GGUF").name == "flux.2-klein"
# FLUX.2-dev is the Mistral-based Flux2Pipeline, a distinct family from klein; its
# gated base repo is reachable with an HF token. It must not collide with klein.
dev = detect_family("unsloth/FLUX.2-dev-GGUF")
assert dev.name == "flux.2-dev"
assert dev.pipeline_class == "Flux2Pipeline"
assert dev.base_repo == "black-forest-labs/FLUX.2-dev"
assert detect_family("black-forest-labs/FLUX.2-dev").name == "flux.2-dev"
# Qwen-Image guides via true_cfg_scale, not guidance_scale.
assert detect_family("unsloth/Qwen-Image-2512-GGUF").cfg_kwarg == "true_cfg_scale"
assert detect_family("unsloth/Z-Image-GGUF").cfg_kwarg == "guidance_scale"
# Qwen-Image-Edit is a SUPPORTED instruction-editing family (its own edit pipeline);
# the most-specific match wins so it doesn't fall back to the generic qwen-image.
edit = detect_family("unsloth/Qwen-Image-Edit-2511-GGUF")
assert edit.name == "qwen-image-edit"
assert edit.pipeline_class == "QwenImageEditPlusPipeline"
assert edit.edit is True
assert detect_family("unsloth/Qwen-Image-Edit-2509-GGUF").name == "qwen-image-edit"
# FLUX Kontext is a SUPPORTED editing family (FluxKontextPipeline); the "kontext"
# keyword is un-rejected for it, and it must win over the generic "flux.1" match.
kontext = detect_family("unsloth/FLUX.1-Kontext-dev-GGUF")
assert kontext.name == "flux.1-kontext"
assert kontext.pipeline_class == "FluxKontextPipeline"
assert kontext.edit is True
assert kontext.cfg_kwarg == "guidance_scale"
# A plain FLUX.1 checkpoint must still resolve to the base flux.1 family, not kontext.
assert detect_family("unsloth/FLUX.1-dev-GGUF").name == "flux.1"
# A plain Qwen-Image checkpoint must still resolve to the base family, not edit.
assert detect_family("unsloth/Qwen-Image-2512-GGUF").name == "qwen-image"
# Krea 2 (diffusers >= 0.39): bf16-only single-stream DiT, no GGUF/sd.cpp mapping.
krea2 = detect_family("krea/Krea-2-Turbo")
assert krea2.name == "krea-2"
assert krea2.pipeline_class == "Krea2Pipeline"
assert krea2.transformer_class == "Krea2Transformer2DModel"
assert krea2.cfg_kwarg == "guidance_scale"
assert krea2.fp16_incompatible is True
assert krea2.sd_cpp_text_encoders == ()
assert detect_family("meta-llama/Llama-3-8B") is None
def test_detect_family_matches_reject_and_alias_by_segment():
# Reject keywords and short aliases must match whole path/name segments, not raw
# substrings, so an unrelated word that merely CONTAINS one does not misroute a
# valid base model (regression: substring matching broke these).
assert detect_family("/models/edited/z-image-turbo-Q4_K_M.gguf").name == "z-image"
assert detect_family("unsloth/Z-Image-Edition-GGUF").name == "z-image"
assert detect_family("/models/kontextual/z-image-turbo-Q4_K_M.gguf").name == "z-image"
# Supported edit families still resolve (edit / kontext are whole tokens there).
assert detect_family("unsloth/Qwen-Image-Edit-2511-GGUF").name == "qwen-image-edit"
assert detect_family("unsloth/FLUX.1-Kontext-dev-GGUF").name == "flux.1-kontext"
# Unsupported variants sharing only a base arch keyword are still rejected.
assert detect_family("unsloth/Qwen-Image-Layered-GGUF") is None
assert detect_family("unsloth/Qwen-Image-2512-Inpaint") is None
def test_detect_family_edit_keyword_scoped_to_basename():
from core.inference.diffusion_families import detect_family_for_pick
# A parent directory named `edit`/`inpaint` must NOT poison a valid pick: only
# the model id / filename basename is scanned for reject keywords. A direct
# local pick arrives as (parent_dir, filename).
assert detect_family("/models/edit") is None # the dir alone is ambiguous
assert detect_family_for_pick("/models/edit", "Z-Image-Turbo-Q4.gguf").name == "z-image"
assert detect_family_for_pick("/models/inpaint", "qwen-image-2512-Q4.gguf").name == "qwen-image"
# A genuinely unsupported variant keyword in the FILENAME still rejects.
assert detect_family_for_pick("/models/misc", "Qwen-Image-Layered-Q4.gguf") is None
def test_detect_family_override():
assert detect_family("local/path", override = "z-image").name == "z-image"
assert detect_family("local/path", override = "zimage").name == "z-image"
assert detect_family("local/path", override = "not-a-family") is None
def test_supported_family_names():
names = supported_family_names()
# The unknown-model error lists these, so the key families must be present.
for expected in ("flux.1", "flux.2-klein", "flux.2-dev", "qwen-image", "z-image", "krea-2"):
assert expected in names
# Every listed name is a valid family_override (round-trips through detect_family).
for name in names:
assert detect_family("some/unknown-repo", override = name) is not None
def test_resolve_base_repo():
fam = detect_family("x", override = "z-image")
assert resolve_base_repo(fam, None) == fam.base_repo
assert resolve_base_repo(fam, " ") == fam.base_repo
assert resolve_base_repo(fam, "custom/base") == "custom/base"
def test_resolve_local_gguf_child(tmp_path):
(tmp_path / "model.gguf").write_bytes(b"x")
assert resolve_local_gguf_child(tmp_path, "model.gguf") == (tmp_path / "model.gguf").resolve()
with pytest.raises(ValueError):
resolve_local_gguf_child(tmp_path, "/etc/passwd")
with pytest.raises(ValueError):
resolve_local_gguf_child(tmp_path, "../secret.gguf")
with pytest.raises(ValueError):
resolve_local_gguf_child(tmp_path, "..\\secret.gguf")
with pytest.raises(FileNotFoundError):
resolve_local_gguf_child(tmp_path, "missing.gguf")
def test_resolve_local_gguf_child_blocks_symlink_escape(tmp_path):
outside = tmp_path / "outside.gguf"
outside.write_bytes(b"secret")
repo = tmp_path / "repo"
repo.mkdir()
try:
(repo / "model.gguf").symlink_to(outside)
except (OSError, NotImplementedError):
pytest.skip("symlinks not supported on this platform")
with pytest.raises(ValueError):
resolve_local_gguf_child(repo, "model.gguf")
# Stubbed runtime for backend lifecycle
class _FakeDtype:
def __init__(self, name: str) -> None:
self._name = name
def __repr__(self) -> str:
return f"torch.{self._name}"
__str__ = __repr__
class _FakeGenerator:
def __init__(self, device = None) -> None:
self.device = device
self.manual = None
def seed(self) -> int:
return 4242
def manual_seed(self, value: int):
self.manual = value
return self
class _FakeImage:
"""Stand-in for a generated PIL image (the route persists it; here we only
count how many come back)."""
class _FakePipe:
def __init__(self) -> None:
self.moved_to = None
self.offloaded = False
self.sequential_offloaded = False
self.vae_tiled = False
self.vae_sliced = False
self.last_kwargs = None
def to(self, device):
self.moved_to = device
return self
def enable_model_cpu_offload(self, device = None) -> None:
self.offloaded = True
self.offload_device = device
def enable_sequential_cpu_offload(self, device = None) -> None:
self.sequential_offloaded = True
self.offload_device = device
def enable_vae_tiling(self) -> None:
self.vae_tiled = True
def enable_vae_slicing(self) -> None:
self.vae_sliced = True
# Explicit signature (not just **kwargs) so generate()'s signature-gated
# guards for negative_prompt / callback_on_step_end actually take effect —
# a **kwargs-only fake would make `"negative_prompt" in signature` always False.
def __call__(
self,
*,
prompt = None,
negative_prompt = None,
callback_on_step_end = None,
guidance_scale = None,
true_cfg_scale = None,
**kwargs,
):
self.last_kwargs = {
"prompt": prompt,
"negative_prompt": negative_prompt,
"callback_on_step_end": callback_on_step_end,
"guidance_scale": guidance_scale,
"true_cfg_scale": true_cfg_scale,
**kwargs,
}
n = kwargs.get("num_images_per_prompt", 1)
return types.SimpleNamespace(images = [_FakeImage() for _ in range(n)])
class _FakePipeline:
last: dict = {}
last_single_file: dict = {}
@classmethod
def from_pretrained(cls, base, **kwargs):
_FakePipeline.last = {"base": base, **kwargs}
return _FakePipe()
@classmethod
def from_single_file(cls, path, **kwargs):
# SDXL-style single-file: the WHOLE pipeline comes from one .safetensors file.
_FakePipeline.last_single_file = {"path": path, **kwargs}
return _FakePipe()
class _FakeTransformer:
last: dict = {}
@classmethod
def from_single_file(cls, path, **kwargs):
_FakeTransformer.last = {"path": path, **kwargs}
return object()
class _FakeImg2ImgPipe:
"""An img2img pipeline call: records the image-conditioned kwargs. Its signature
declares image/strength but NOT width/height, mirroring real img2img pipelines
(which derive the output size from the input image)."""
last_kwargs: dict = {}
def __call__(
self,
*,
prompt = None,
image = None,
strength = None,
negative_prompt = None,
callback_on_step_end = None,
guidance_scale = None,
true_cfg_scale = None,
**kwargs,
):
_FakeImg2ImgPipe.last_kwargs = {
"prompt": prompt,
"image": image,
"strength": strength,
**kwargs,
}
n = kwargs.get("num_images_per_prompt", 1)
return types.SimpleNamespace(images = [_FakeImage() for _ in range(n)])
class _FakeImg2ImgPipeline:
built_from: object = None
from_pipe_kwargs: dict = {}
@classmethod
def from_pipe(cls, base_pipe, **kwargs):
_FakeImg2ImgPipeline.built_from = base_pipe
_FakeImg2ImgPipeline.from_pipe_kwargs = kwargs
return _FakeImg2ImgPipe()
class _FakeInpaintPipe:
"""An inpaint pipeline call: records image + mask_image + strength. Real inpaint
pipelines take both an init image and a grayscale mask and derive output size from
the input, so width/height are not in its signature."""
last_kwargs: dict = {}
def __call__(
self,
*,
prompt = None,
image = None,
mask_image = None,
strength = None,
negative_prompt = None,
callback_on_step_end = None,
guidance_scale = None,
true_cfg_scale = None,
**kwargs,
):
_FakeInpaintPipe.last_kwargs = {
"prompt": prompt,
"image": image,
"mask_image": mask_image,
"strength": strength,
**kwargs,
}
n = kwargs.get("num_images_per_prompt", 1)
return types.SimpleNamespace(images = [_FakeImage() for _ in range(n)])
class _FakeInpaintPipeline:
built_from: object = None
@classmethod
def from_pipe(cls, base_pipe, **kwargs):
_FakeInpaintPipeline.built_from = base_pipe
return _FakeInpaintPipe()
@pytest.fixture
def fake_runtime(monkeypatch):
torch = types.ModuleType("torch")
torch.bfloat16 = _FakeDtype("bfloat16")
torch.float16 = _FakeDtype("float16")
torch.float32 = _FakeDtype("float32")
torch.Generator = _FakeGenerator
torch.cuda = types.SimpleNamespace(is_available = lambda: False)
torch.backends = types.SimpleNamespace(mps = None)
# generate() wraps the pipe call in torch.inference_mode(); a no-op CM here.
torch.inference_mode = lambda: contextlib.nullcontext()
diffusers = types.ModuleType("diffusers")
diffusers.GGUFQuantizationConfig = lambda compute_dtype = None: ("quant", compute_dtype)
diffusers.ZImagePipeline = _FakePipeline
diffusers.ZImageTransformer2DModel = _FakeTransformer
diffusers.ZImageImg2ImgPipeline = _FakeImg2ImgPipeline
diffusers.ZImageInpaintPipeline = _FakeInpaintPipeline
# Qwen-Image too, so the true_cfg_scale cfg-kwarg path is exercisable.
diffusers.QwenImagePipeline = _FakePipeline
diffusers.QwenImageTransformer2DModel = _FakeTransformer
diffusers.QwenImageImg2ImgPipeline = _FakeImg2ImgPipeline
diffusers.QwenImageInpaintPipeline = _FakeInpaintPipeline
# Instruction-editing pipeline (Qwen-Image-Edit): its own pipeline IS the loaded one.
diffusers.QwenImageEditPlusPipeline = _FakePipeline
# Ideogram 4, so its guidance_scale/guidance_schedule pairing is exercisable. It loads
# only as a full pipeline (two DiTs), assembled per-component by load_ideogram4_pipeline
# -- stub that to a fake pipe so the guidance path is reachable without real weights.
diffusers.Ideogram4Pipeline = _FakePipeline
diffusers.Ideogram4Transformer2DModel = _FakeTransformer
# SDXL: a U-Net family. Its single-file checkpoint is the whole pipeline, so the
# pipeline class carries from_single_file; UNet2DConditionModel is the denoiser
# class (fetched but unused on the pipeline/single-file-pipeline paths).
diffusers.StableDiffusionXLPipeline = _FakePipeline
diffusers.UNet2DConditionModel = _FakeTransformer
diffusers.StableDiffusionXLImg2ImgPipeline = _FakeImg2ImgPipeline
diffusers.StableDiffusionXLInpaintPipeline = _FakeInpaintPipeline
monkeypatch.setattr(
"core.inference.diffusion.load_ideogram4_pipeline",
lambda repo_id, dtype, hf_token = None: _FakePipe(),
)
monkeypatch.setitem(sys.modules, "torch", torch)
monkeypatch.setitem(sys.modules, "diffusers", diffusers)
# 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)
_FakePipeline.last = {}
_FakePipeline.last_single_file = {}
_FakeTransformer.last = {}
_FakeImg2ImgPipeline.built_from = None
_FakeImg2ImgPipe.last_kwargs = {}
_FakeInpaintPipeline.built_from = None
_FakeInpaintPipe.last_kwargs = {}
yield
def test_load_generate_unload_gguf(fake_runtime, tmp_path):
(tmp_path / "model.gguf").write_bytes(b"weights")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
base_repo = "base/repo",
family_override = "z-image",
hf_token = "hf_secret",
)
assert status["loaded"] is True
assert status["family"] == "z-image"
assert status["base_repo"] == "base/repo"
assert status["device"] == "cpu"
assert status["dtype"] == "float32"
assert status["cpu_offload"] is False
# Transformer built from the local GGUF, pipeline assembled from the base repo.
assert _FakeTransformer.last["path"] == str((tmp_path / "model.gguf").resolve())
assert _FakeTransformer.last["subfolder"] == "transformer"
# The token reaches the (possibly gated) base config fetch and the pipeline.
assert _FakeTransformer.last["token"] == "hf_secret"
assert _FakePipeline.last["base"] == "base/repo"
assert "transformer" in _FakePipeline.last
gen = backend.generate(
prompt = "a sloth", negative_prompt = "blurry", width = 512, height = 512, steps = 4, guidance = 3.0
)
assert gen["seed"] == 4242 # random seed reported back
assert gen["repo_id"] == str(tmp_path) # echoed so the route can record the model
assert len(gen["images"]) == 1 # PIL images handed to the route for persistence
# z-image guides via guidance_scale (not true_cfg_scale); the signature-gated
# negative_prompt and per-step callback both reach the pipeline call.
call = backend._state.pipe.last_kwargs
assert call["guidance_scale"] == 3.0 and call["true_cfg_scale"] is None
assert call["negative_prompt"] == "blurry"
assert callable(call["callback_on_step_end"])
gen2 = backend.generate(prompt = "again", seed = 99)
assert gen2["seed"] == 99
# batch_size produces that many images in one call, all sharing the seed.
batch = backend.generate(prompt = "batch", seed = 7, batch_size = 3)
assert len(batch["images"]) == 3 and batch["seed"] == 7
assert backend.unload()["loaded"] is False
assert backend.is_loaded is False
def test_dense_speed_auto_defers_compile_to_third_generation(fake_runtime, tmp_path, monkeypatch):
# Dense models with speed unset stay bit-identical eager for the first two
# generations; the 3rd engages the `default` profile mid-session (repeated
# use amortises the one-time compile), upgrading attention alongside it.
from core.inference import diffusion as dmod
monkeypatch.setattr(dmod, "compile_eligible", lambda *a, **k: True)
monkeypatch.setattr(
dmod,
"apply_speed_optims",
lambda pipe, target, **k: {"compiled": k.get("speed_mode") == "default"},
)
monkeypatch.setattr(dmod, "apply_attention_backend", lambda pipe, backend, logger = None: backend)
monkeypatch.setattr(
dmod,
"select_attention_backend",
lambda target, requested, speed_active = False: ("_native_cudnn" if speed_active else None),
)
monkeypatch.setattr(dmod.compile_cache, "begin", lambda **k: None)
(tmp_path / "model.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
)
assert status["speed_mode"] == "off"
assert status["resolved"]["speed_mode"]["value"] == "deferred"
assert status["resolved"]["speed_mode"]["source"] == "auto"
backend.generate(prompt = "one")
backend.generate(prompt = "two")
assert backend.status()["speed_mode"] == "off" # first two stay exact eager
backend.generate(prompt = "three")
status3 = backend.status()
assert status3["speed_mode"] == "default"
assert "compiled" in status3["speed_optims"]
assert status3["attention_backend"] == "_native_cudnn"
assert status3["resolved"]["speed_mode"]["value"] == "default"
# An explicit "off" is pinned: no deferral, still eager after 3 generations.
backend.unload()
status_off = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
speed_mode = "off",
)
assert status_off["resolved"]["speed_mode"]["value"] == "off"
for p in ("a", "b", "c"):
backend.generate(prompt = p)
assert backend.status()["speed_mode"] == "off"
backend.unload()
def test_deferred_speed_skips_when_lora_requested(fake_runtime, tmp_path, monkeypatch):
# A compiled transformer rejects LoRA (supports_lora is False once compiled), and _apply_loras
# raises before its unchanged-selection no-op, so engaging the deferred compile on a generation
# that requests a LoRA would permanently break every LoRA generation on this load. The deferral
# must skip while a LoRA is requested and engage only on a later LoRA-free generation.
from core.inference import diffusion as dmod
monkeypatch.setattr(dmod, "compile_eligible", lambda *a, **k: True)
engaged: list = []
def fake_engage(self, state):
engaged.append(state.generation_count)
state.speed_deferred = False # mirror the real helper: engage once, then clear
monkeypatch.setattr(DiffusionBackend, "_engage_deferred_speed", fake_engage)
# LoRA loading is covered elsewhere; stub it so this test needs no adapter file.
monkeypatch.setattr(DiffusionBackend, "_apply_loras", lambda self, state, loras, cancel: None)
(tmp_path / "model.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
)
backend.generate(prompt = "one")
backend.generate(prompt = "two")
# 3rd generation requests a LoRA: the deferral must be skipped (pipe stays eager, LoRA-capable).
backend.generate(prompt = "three", loras = [("adapter", 1.0)])
assert engaged == []
# 4th generation without a LoRA: the deferral now engages (the guard is LoRA-specific, not off).
backend.generate(prompt = "four")
assert len(engaged) == 1
def test_deferred_speed_skips_while_adapter_attached(fake_runtime, tmp_path, monkeypatch):
# Even a generation that requests NO LoRA must defer the compile while an adapter from a PRIOR
# generation is still attached: _apply_loras runs AFTER the engage, so compiling here would bake
# the resident adapter into the graph and the subsequent unload (swallowed on a compiled pipe)
# would leave it active forever -- silent wrong output. Defer until _apply_loras clears it.
from core.inference import diffusion as dmod
monkeypatch.setattr(dmod, "compile_eligible", lambda *a, **k: True)
engaged: list = []
def fake_engage(self, state):
engaged.append(state.generation_count)
state.speed_deferred = False
monkeypatch.setattr(DiffusionBackend, "_engage_deferred_speed", fake_engage)
# Track the attached set on the pipe, mirroring the real _apply_loras marker (_unsloth_loras).
def fake_apply(self, state, loras, cancel):
specs = [(i, w) for (i, w) in (loras or []) if w != 0]
state.pipe._unsloth_loras = tuple(specs)
monkeypatch.setattr(DiffusionBackend, "_apply_loras", fake_apply)
(tmp_path / "model.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
)
# Gens 1-2 attach an adapter, so it is still resident going into gen 3.
backend.generate(prompt = "one", loras = [("adapter", 1.0)])
backend.generate(prompt = "two", loras = [("adapter", 1.0)])
# Gen 3 requests NO LoRA but the adapter is still attached -> defer (no compile-with-adapter).
backend.generate(prompt = "three")
assert engaged == []
# Gen 3's _apply_loras([]) cleared the adapter; gen 4 is genuinely LoRA-free -> engage.
backend.generate(prompt = "four")
assert len(engaged) == 1
def test_deferred_speed_preserves_explicit_attention(fake_runtime, tmp_path, monkeypatch):
# A dense model loaded with Speed left on Auto but Attention explicitly pinned
# (e.g. "native" to avoid cuDNN) must KEEP that choice when the 3rd generation
# engages the deferred `default` profile. The auto cuDNN upgrade only applies when
# attention was left on auto -- never when the caller pinned a backend.
from core.inference import diffusion as dmod
monkeypatch.setattr(dmod, "compile_eligible", lambda *a, **k: True)
monkeypatch.setattr(
dmod,
"apply_speed_optims",
lambda pipe, target, **k: {"compiled": k.get("speed_mode") == "default"},
)
monkeypatch.setattr(dmod, "apply_attention_backend", lambda pipe, backend, logger = None: backend)
# A select mock that -- unlike a bare "auto -> cuDNN" stub -- HONORS an explicit
# request: "native" stays on the default (None) even under a speed profile, and only
# a left-unset ("auto"/None) request upgrades to cuDNN when speed is active.
def fake_select(
target,
requested,
speed_active = False,
):
if requested in (None, "", "auto"):
return "_native_cudnn" if speed_active else None
if str(requested).lower() in ("native", "sdpa"):
return None
return requested
monkeypatch.setattr(dmod, "select_attention_backend", fake_select)
monkeypatch.setattr(dmod.compile_cache, "begin", lambda **k: None)
(tmp_path / "model.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
attention_backend = "native",
)
backend.generate(prompt = "one")
backend.generate(prompt = "two")
backend.generate(prompt = "three") # deferred profile engages here
status = backend.status()
assert status["speed_mode"] == "default" # the compile profile still engaged
assert "compiled" in status["speed_optims"]
# The pinned "native" survived: NOT silently upgraded to cuDNN.
assert status["attention_backend"] is None
assert status["resolved"]["attention_backend"]["value"] == "native"
assert status["resolved"]["attention_backend"]["source"] == "explicit"
# Control: with attention left on auto, the same 3rd-generation deferral DOES upgrade
# to cuDNN -- so the assertion above is not vacuously passing.
backend.unload()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
)
for p in ("a", "b", "c"):
backend.generate(prompt = p)
assert backend.status()["attention_backend"] == "_native_cudnn"
backend.unload()
def _tiny_png_b64() -> str:
import base64
import io
from PIL import Image
buf = io.BytesIO()
Image.new("RGB", (64, 64), (120, 30, 30)).save(buf, format = "PNG")
return base64.b64encode(buf.getvalue()).decode()
def test_generate_img2img_uses_from_pipe(fake_runtime, tmp_path):
"""An init_image routes generate() through the family's img2img pipeline, built via
Pipeline.from_pipe around the loaded pipe (no reload), with image + strength passed
and width/height dropped (the img2img pipe derives size from the input image)."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
# The loaded family advertises the image-conditioned workflows for UI gating
# (upscale rides the img2img pipeline, so it appears whenever img2img does).
assert backend.status()["workflows"] == ["txt2img", "img2img", "upscale", "inpaint", "outpaint"]
loaded_pipe = backend._state.pipe
out = backend.generate(
prompt = "a car at sunset",
steps = 4,
guidance = 0.0,
seed = 3,
init_image = _tiny_png_b64(),
strength = 0.5,
)
assert len(out["images"]) == 1
# from_pipe was handed the loaded text-to-image pipe (component reuse, no reload).
assert _FakeImg2ImgPipeline.built_from is loaded_pipe
# ...and with torch_dtype=None so from_pipe SKIPS its default float32 recast, which
# both upcasts the reused bf16 modules and crashes on torchao-quantized weights.
assert _FakeImg2ImgPipeline.from_pipe_kwargs.get("torch_dtype", "MISSING") is None
call = _FakeImg2ImgPipe.last_kwargs
assert call["image"] is not None # decoded source image passed through
assert call["strength"] == 0.5
assert "width" not in call and "height" not in call # img2img derives size from image
# A txt2img call after it still uses the base pipe (no image kwarg).
backend.generate(prompt = "plain", steps = 4, seed = 1)
assert backend._state.pipe.last_kwargs.get("image") is None
def test_generate_img2img_unsupported_family_raises(fake_runtime, tmp_path, monkeypatch):
"""A family with no image-conditioning at all (no img2img/inpaint/edit/reference) rejects
an init_image with a clear error rather than failing deep in the pipeline."""
from core.inference.diffusion_families import DiffusionFamily
# A synthetic txt2img-only family: no img2img/inpaint pipeline, not edit, not reference.
# (Every shipped family now supports some image workflow, so build one for this case.)
plain = DiffusionFamily(
name = "plain-test",
pipeline_class = "ZImagePipeline",
transformer_class = "ZImageTransformer2DModel",
base_repo = "base/repo",
)
monkeypatch.setattr(
"core.inference.diffusion.detect_family_for_pick",
lambda repo_id, gguf_filename = None, override = None: plain,
)
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo")
assert backend.status()["workflows"] == ["txt2img"]
with pytest.raises(ValueError, match = "img2img"):
backend.generate(prompt = "x", steps = 4, init_image = _tiny_png_b64())
def test_generate_rejects_conditioning_without_init_image(fake_runtime, tmp_path):
"""mask / upscale / reference all need an input image; without one they must raise a
clear ValueError rather than silently degrading to txt2img."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
with pytest.raises(ValueError, match = "mask_image requires"):
backend.generate(prompt = "x", steps = 4, mask_image = _mask_b64(64))
with pytest.raises(ValueError, match = "upscale requires"):
backend.generate(prompt = "x", steps = 4, upscale = 2.0)
with pytest.raises(ValueError, match = "reference_images require"):
backend.generate(prompt = "x", steps = 4, reference_images = [_tiny_png_b64()])
def test_generate_rejects_reference_on_unsupported_family(fake_runtime, tmp_path):
"""A non-reference family rejects reference_images instead of silently dropping them."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
with pytest.raises(ValueError, match = "Reference images are not supported"):
backend.generate(
prompt = "x",
steps = 4,
init_image = _tiny_png_b64(),
reference_images = [_tiny_png_b64()],
)
def test_generate_upscale_enlarges_and_low_strength(fake_runtime, tmp_path):
"""An init_image + upscale factor routes generate() through the family's img2img
pipeline (hires fix): the source is enlarged to size*factor (rounded to /16) before the
denoise, the strength defaults low, and the factor is capped so a huge value can't OOM."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
# Upscale rides the img2img pipeline, so it is advertised alongside img2img.
assert "upscale" in backend.status()["workflows"]
loaded_pipe = backend._state.pipe
out = backend.generate(
prompt = "a crisp photo",
steps = 4,
guidance = 0.0,
seed = 3,
init_image = _tiny_png_b64(),
upscale = 2.0, # 64 -> 128, no explicit strength
)
assert len(out["images"]) == 1
# Reuses the resident modules via from_pipe (no reload, no extra VRAM).
assert _FakeImg2ImgPipeline.built_from is loaded_pipe
call = _FakeImg2ImgPipe.last_kwargs
# The image handed to the pipe is the ENLARGED source (64 * 2 = 128, already /16).
assert call["image"].size == (128, 128)
# Strength defaults to the hires-fix value when the caller sends none.
assert call["strength"] == 0.35
# The factor is capped at 4x so a large request can't blow up the VAE/transformer.
backend.generate(
prompt = "x",
steps = 4,
seed = 1,
init_image = _tiny_png_b64(),
upscale = 99.0,
)
assert _FakeImg2ImgPipe.last_kwargs["image"].size == (256, 256) # 64 * 4 (capped)
# An explicit strength overrides the hires-fix default.
backend.generate(
prompt = "x",
steps = 4,
seed = 1,
init_image = _tiny_png_b64(),
upscale = 1.5,
strength = 0.2,
)
assert _FakeImg2ImgPipe.last_kwargs["strength"] == 0.2
# 64 * 1.5 = 96, already a multiple of 16.
assert _FakeImg2ImgPipe.last_kwargs["image"].size == (96, 96)
def _png_b64(side: int) -> str:
import base64
import io
from PIL import Image
buf = io.BytesIO()
Image.new("RGB", (side, side), (10, 20, 30)).save(buf, format = "PNG")
return base64.b64encode(buf.getvalue()).decode()
def test_decode_image_rejects_oversized(fake_runtime, tmp_path):
"""An input image larger than the per-side cap is rejected with a clear error (protects
img2img / inpaint / reference from decompression-bomb / OOM inputs), not a 500."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
with pytest.raises(ValueError, match = "too large"):
backend.generate(prompt = "x", steps = 4, init_image = _png_b64(4112)) # > 4096/side
def test_upscale_output_is_capped(fake_runtime, tmp_path):
"""Upscale bounds the absolute output side to 2048 even when input*factor exceeds it, so a
large upload at 4x can't OOM the VAE/transformer."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
backend.generate(prompt = "x", steps = 4, seed = 1, init_image = _png_b64(1024), upscale = 4.0)
# 1024 * 4 = 4096 -> clamped to 2048 (longest side), still a multiple of 16.
assert _FakeImg2ImgPipe.last_kwargs["image"].size == (2048, 2048)
def _mask_b64(side: int) -> str:
import base64
import io
from PIL import Image
buf = io.BytesIO()
img = Image.new("L", (side, side), 0)
for y in range(side // 4, 3 * side // 4):
for x in range(side // 4, 3 * side // 4):
img.putpixel((x, y), 255)
img.save(buf, format = "PNG")
return base64.b64encode(buf.getvalue()).decode()
def test_img2img_snaps_non_multiple_of_16(fake_runtime, tmp_path):
"""An odd-sized img2img upload (not divisible by 16) is auto-resized to the nearest
multiple of 16 so the pipeline's divisibility check passes instead of erroring."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
backend.generate(prompt = "x", steps = 4, seed = 1, init_image = _png_b64(186), strength = 0.5)
# 186 / 16 = 11.625 -> round to 12 -> 192.
assert _FakeImg2ImgPipe.last_kwargs["image"].size == (192, 192)
def test_inpaint_snaps_image_and_mask_together(fake_runtime, tmp_path):
"""Inpaint snaps the odd-sized input to /16 AND resizes the mask to match, so the image
and mask stay aligned (a mismatch would crash the inpaint pipeline)."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
backend.generate(
prompt = "x",
steps = 4,
seed = 1,
init_image = _png_b64(186),
mask_image = _mask_b64(186),
strength = 0.5,
)
assert _FakeInpaintPipe.last_kwargs["image"].size == (192, 192)
assert _FakeInpaintPipe.last_kwargs["mask_image"].size == (192, 192)
def test_generate_reference_uses_loaded_pipe_at_slider_size(fake_runtime, tmp_path):
"""A reference family (FLUX.2-klein) advertises txt2img + reference, and a generate with
an init_image passes it as the loaded pipe's `image` arg (no from_pipe, no strength) while
the output size stays the REQUESTED slider size (the pipe resizes the reference itself)."""
import diffusers
diffusers.Flux2KleinPipeline = _FakePipeline
diffusers.Flux2KleinInpaintPipeline = _FakeInpaintPipeline
diffusers.Flux2Transformer2DModel = _FakeTransformer
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
base_repo = "base/repo",
family_override = "flux.2-klein",
)
# FLUX.2-klein: txt2img + reference (own pipe) + inpaint (dedicated pipe). No img2img class,
# so no img2img/upscale.
assert backend.status()["workflows"] == ["txt2img", "reference", "inpaint"]
loaded_pipe = backend._state.pipe
out = backend.generate(
prompt = "a portrait in this style",
steps = 6,
guidance = 4.0,
seed = 5,
width = 768,
height = 512,
init_image = _tiny_png_b64(),
strength = 0.5,
)
assert len(out["images"]) == 1
call = loaded_pipe.last_kwargs
assert call["image"] is not None # reference handed to the loaded pipe
assert call["width"] == 768 and call["height"] == 512 # OUTPUT size = sliders, not input
assert "strength" not in call # reference conditioning has no strength
assert "mask_image" not in call
# Guidance flows via guidance_scale (FLUX.2 default behaviour).
assert call["guidance_scale"] == 4.0
# Multi-reference: extra reference_images are combined with init_image into a LIST so the
# model can blend several references (subject + style).
backend.generate(
prompt = "combine these",
steps = 6,
seed = 9,
width = 1024,
height = 1024,
init_image = _tiny_png_b64(),
reference_images = [_tiny_png_b64(), _tiny_png_b64()],
)
img_arg = loaded_pipe.last_kwargs["image"]
assert isinstance(img_arg, list) and len(img_arg) == 3 # primary + 2 extras
# Branch ordering: an init image + MASK on a reference family must route to inpaint (the
# dedicated pipeline), NOT be swallowed by the reference branch (which ignores the mask).
backend.generate(
prompt = "repaint here",
steps = 6,
seed = 2,
init_image = _tiny_png_b64(),
mask_image = _tiny_mask_b64(),
strength = 0.8,
)
assert _FakeInpaintPipeline.built_from is loaded_pipe # built via from_pipe off the load
assert _FakeInpaintPipe.last_kwargs["mask_image"] is not None
assert _FakeInpaintPipe.last_kwargs["strength"] == 0.8
# Without an init image the same family does plain txt2img (no image arg).
backend.generate(prompt = "just text", steps = 6, seed = 1)
assert backend._state.pipe.last_kwargs.get("image") is None
def _tiny_mask_b64() -> str:
import base64
import io
from PIL import Image
buf = io.BytesIO()
# A grayscale mask: white square (repaint) on black (keep).
img = Image.new("L", (64, 64), 0)
for y in range(16, 48):
for x in range(16, 48):
img.putpixel((x, y), 255)
img.save(buf, format = "PNG")
return base64.b64encode(buf.getvalue()).decode()
def test_generate_inpaint_uses_from_pipe(fake_runtime, tmp_path):
"""An init_image + mask_image routes generate() through the family's inpaint pipeline,
built via Pipeline.from_pipe around the loaded pipe (no reload), with the decoded image
+ mask + strength passed through and width/height dropped (size derives from the input)."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
loaded_pipe = backend._state.pipe
out = backend.generate(
prompt = "a red door",
steps = 4,
guidance = 0.0,
seed = 5,
init_image = _tiny_png_b64(),
mask_image = _tiny_mask_b64(),
strength = 0.7,
)
assert len(out["images"]) == 1
# The inpaint pipe (not img2img) was selected and built from the loaded pipe.
assert _FakeInpaintPipeline.built_from is loaded_pipe
assert _FakeImg2ImgPipeline.built_from is None
call = _FakeInpaintPipe.last_kwargs
assert call["image"] is not None and call["mask_image"] is not None
assert call["strength"] == 0.7
assert "width" not in call and "height" not in call # inpaint derives size from image
def test_image_conditioned_passes_image_size_not_slider(fake_runtime, tmp_path):
"""When the workflow pipe DOES accept width/height, an image-conditioned call must pass
the INPUT IMAGE's size, never the txt2img slider size -- otherwise a non-slider-sized
input (e.g. a 1536px outpaint canvas with a 1024 slider) mismatches the latents
("tensor a (128) must match tensor b (192)"). Covers Transform + Extend with any size."""
import base64
import io
from PIL import Image
class _SizePipe:
last: dict = {}
def __call__(
self,
*,
prompt = None,
image = None,
strength = None,
width = None,
height = None,
negative_prompt = None,
callback_on_step_end = None,
guidance_scale = None,
true_cfg_scale = None,
**kwargs,
):
_SizePipe.last = {"width": width, "height": height}
n = kwargs.get("num_images_per_prompt", 1)
return types.SimpleNamespace(images = [_FakeImage() for _ in range(n)])
class _SizePipeline:
@classmethod
def from_pipe(cls, base_pipe, **kwargs):
return _SizePipe()
import diffusers
diffusers.ZImageImg2ImgPipeline = _SizePipeline
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
buf = io.BytesIO()
Image.new("RGB", (96, 64), (10, 20, 30)).save(buf, format = "PNG") # non-square, non-slider
b64 = base64.b64encode(buf.getvalue()).decode()
backend.generate(prompt = "x", steps = 4, width = 1024, height = 1024, init_image = b64, strength = 0.5)
# The pipe got the IMAGE's 96x64, not the 1024x1024 slider.
assert _SizePipe.last == {"width": 96, "height": 64}
def test_edit_family_uses_own_pipeline_and_requires_image(fake_runtime, tmp_path):
"""An instruction-editing family (Qwen-Image-Edit) exposes only the 'edit' workflow,
runs the image through its OWN loaded pipeline (no from_pipe), and rejects a call with
no input image."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
base_repo = "Qwen/Qwen-Image-Edit-2511",
family_override = "qwen-image-edit",
)
# Edit families advertise only the edit workflow (no txt2img / img2img / inpaint).
assert backend.status()["workflows"] == ["edit"]
loaded_pipe = backend._state.pipe
out = backend.generate(
prompt = "make it night",
steps = 8,
guidance = 4.0,
seed = 1,
init_image = _tiny_png_b64(),
)
assert len(out["images"]) == 1
# The loaded pipe handled it directly -- no from_pipe img2img/inpaint was built.
assert backend._state.pipe is loaded_pipe
assert _FakeImg2ImgPipeline.built_from is None and _FakeInpaintPipeline.built_from is None
assert loaded_pipe.last_kwargs.get("image") is not None
# An edit model with no input image fails fast with a clear message.
with pytest.raises(ValueError, match = "image"):
backend.generate(prompt = "make it night", steps = 8)
def test_load_pipeline_kind_uses_from_pretrained(fake_runtime):
"""A full-pipeline (no single-file) load on an unsloth/* repo builds the pipe with
pipeline_cls.from_pretrained(repo_id) -- NO single-file transformer build, NO GGUF
quant config -- so an embedded bnb-4bit config is reloaded by diffusers itself."""
backend = DiffusionBackend()
status = backend.load_pipeline(
"unsloth/Z-Image-Turbo-unsloth-bnb-4bit", family_override = "z-image"
)
assert status["loaded"] is True
assert status["family"] == "z-image"
# from_pretrained pointed at the repo itself (it IS its own base), with no transformer.
assert _FakePipeline.last["base"] == "unsloth/Z-Image-Turbo-unsloth-bnb-4bit"
assert "transformer" not in _FakePipeline.last
# The GGUF single-file build path was never taken.
assert _FakeTransformer.last == {}
def test_load_single_file_safetensors_no_gguf_config(fake_runtime, tmp_path):
"""A single-file *.safetensors transformer is built with from_single_file WITHOUT the
GGUF dequant config (it carries its own dtype), then assembled from the base repo."""
(tmp_path / "model.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
base_repo = "base/repo",
family_override = "qwen-image",
)
assert status["loaded"] is True
assert _FakeTransformer.last["path"] == str((tmp_path / "model.safetensors").resolve())
assert _FakeTransformer.last["subfolder"] == "transformer"
# No GGUF quant config on the safetensors path (the GGUF path sets one).
assert "quantization_config" not in _FakeTransformer.last
assert _FakePipeline.last["base"] == "base/repo"
assert "transformer" in _FakePipeline.last
def test_load_sdxl_pipeline_from_pretrained(fake_runtime):
"""SDXL as a full pipeline (no single-file name) loads via pipeline_cls.from_pretrained
on the allowlisted official base repo -- no U-Net single-file build, no GGUF config.
A U-Net family must NOT try to build a transformer from a single file."""
backend = DiffusionBackend()
status = backend.load_pipeline("stabilityai/stable-diffusion-xl-base-1.0")
assert status["loaded"] is True
assert status["family"] == "sdxl"
assert _FakePipeline.last["base"] == "stabilityai/stable-diffusion-xl-base-1.0"
assert "transformer" not in _FakePipeline.last
# Neither single-file path (transformer-only nor whole-pipeline) was taken.
assert _FakeTransformer.last == {}
assert _FakePipeline.last_single_file == {}
def test_load_sdxl_single_file_uses_pipeline_from_single_file(fake_runtime, tmp_path):
"""A single-file SDXL *.safetensors is the WHOLE pipeline: it must load via
pipeline_cls.from_single_file(path, config=base), NOT transformer_cls.from_single_file
(UNet2DConditionModel has no companion-transformer assembly here)."""
(tmp_path / "sdxl.safetensors").write_bytes(b"weights")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "sdxl.safetensors", family_override = "sdxl"
)
assert status["loaded"] is True
assert status["family"] == "sdxl"
# The whole-pipeline single-file path was taken with the base repo as config.
assert _FakePipeline.last_single_file["path"] == str((tmp_path / "sdxl.safetensors").resolve())
assert _FakePipeline.last_single_file["config"] == "stabilityai/stable-diffusion-xl-base-1.0"
# The transformer-only single-file build was NOT taken.
assert _FakeTransformer.last == {}
def test_load_sdxl_allowlisted_turbo_repo_is_trusted(fake_runtime):
"""The official sdxl-turbo repo is on the non-GGUF allowlist, so a full-pipeline load
is permitted even though it is not under unsloth/*."""
backend = DiffusionBackend()
status = backend.load_pipeline("stabilityai/sdxl-turbo")
assert status["loaded"] is True
assert status["family"] == "sdxl"
def test_load_pipeline_rejects_non_unsloth_repo(fake_runtime):
backend = DiffusionBackend()
with pytest.raises(ValueError, match = "unsloth"):
backend.load_pipeline("randomorg/Z-Image-bnb-4bit", family_override = "z-image")
def test_load_sdxl_rejects_untrusted_repo(fake_runtime):
"""A random non-allowlisted, non-unsloth repo is still rejected for a full pipeline
load even when it detects as SDXL -- the allowlist is exact-match only."""
backend = DiffusionBackend()
with pytest.raises(ValueError, match = "unsloth"):
backend.load_pipeline("randomorg/my-sdxl-merge", family_override = "sdxl")
def test_detect_family_rejects_layered():
# Qwen-Image-Layered needs a dedicated pipeline (additional_t_cond); it must be
# rejected so it fails fast at load instead of crashing at the first denoise step.
assert detect_family("unsloth/Qwen-Image-Layered-GGUF") is None
assert detect_family("unsloth/qwen_image_layered") is None
def test_failed_load_rolls_back_eager_patches(fake_runtime, tmp_path, monkeypatch):
"""A load failure AFTER the eager patches install but BEFORE the _LoadState commit must
roll the process-wide patches back, so the next bit-identical `off` load is not
contaminated (the asymmetric-cleanup bug the reviewers flagged)."""
from core.inference import diffusion as diff_mod
from core.inference import diffusion_eager_patches as ep
(tmp_path / "model.gguf").write_bytes(b"x")
ep.uninstall_patches() # clean slate
def _boom(*_a, **_k):
raise RuntimeError("placement boom")
# apply_memory_plan runs AFTER the patches are installed, before _LoadState commits.
monkeypatch.setattr(diff_mod, "apply_memory_plan", _boom)
backend = DiffusionBackend()
with pytest.raises(RuntimeError):
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
family_override = "z-image",
base_repo = "base/repo",
speed_mode = "eager", # != off -> installs the shared patches
)
assert ep.is_installed() is False # rolled back by the load-failure finally
assert backend.is_loaded is False
def test_cpu_offload_ignored_off_cuda(fake_runtime, tmp_path):
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
family_override = "z-image",
base_repo = "base/repo",
cpu_offload = True,
)
# No CUDA in the stub, so offload is not engaged.
assert status["cpu_offload"] is False
def test_low_vram_ignored_off_cuda(fake_runtime, tmp_path):
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
family_override = "z-image",
base_repo = "base/repo",
memory_mode = "low_vram",
)
# No CUDA in the stub, so offload is not engaged regardless of the request.
assert status["cpu_offload"] is False
def test_generate_without_load_raises(fake_runtime):
backend = DiffusionBackend()
with pytest.raises(RuntimeError):
backend.generate(prompt = "x")
def test_failed_load_restores_backend_flags(fake_runtime, tmp_path, monkeypatch):
# A failure AFTER apply_speed_optims (here an OOM in apply_memory_plan) must go
# through the load's try/finally and restore the process-global TF32 / cudnn flags,
# so a later `off` load is still bit-identical, and must not commit a partial state.
# Regression: a refactor dropped this guard, leaking the flags on a failed load.
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
restored: list = []
cleared: list = []
monkeypatch.setattr(
"core.inference.diffusion.restore_backend_flags", lambda snap: restored.append(snap)
)
monkeypatch.setattr("core.inference.diffusion.clear_gpu_cache", lambda: cleared.append(True))
monkeypatch.setattr(
"core.inference.diffusion.apply_memory_plan",
lambda *a, **k: (_ for _ in ()).throw(RuntimeError("CUDA out of memory")),
)
with pytest.raises(RuntimeError, match = "out of memory"):
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
family_override = "z-image",
base_repo = "base/repo",
speed_mode = "max",
)
assert restored, "restore_backend_flags was not called on the failed-load path"
assert cleared, "clear_gpu_cache was not called on the failed-load path (VRAM leak)"
assert backend._state is None and backend.is_loaded is False
def test_resolve_base_repo_prefers_caller_then_hf_tag_then_fallback(monkeypatch):
from core.inference import diffusion
from core.inference.diffusion_families import detect_family
fam = detect_family("unsloth/Qwen-Image-2512-GGUF")
monkeypatch.setattr(diffusion, "_hf_base_model", lambda repo, tok: "Qwen/Qwen-Image-2512")
# Caller's explicit base wins and the HF tag is not consulted.
assert (
diffusion._resolve_base_repo("unsloth/Qwen-Image-2512-GGUF", "my/base", fam, None)
== "my/base"
)
# No caller base: the repo's base_model tag (the variant base) is used.
assert (
diffusion._resolve_base_repo("unsloth/Qwen-Image-2512-GGUF", None, fam, None)
== "Qwen/Qwen-Image-2512"
)
# No caller base and no tag: the family fallback.
monkeypatch.setattr(diffusion, "_hf_base_model", lambda repo, tok: None)
assert (
diffusion._resolve_base_repo("unsloth/Qwen-Image-2512-GGUF", " ", fam, None)
== fam.base_repo
)
def test_load_without_gguf_raises():
backend = DiffusionBackend()
# No gguf_filename -> a full-pipeline load, gated to unsloth/*; a non-unsloth repo
# is rejected before any GPU/network work.
with pytest.raises(ValueError, match = "unsloth"):
backend.load_pipeline("some-org/Z-Image-bnb-4bit")
def test_load_unknown_family_raises():
backend = DiffusionBackend()
with pytest.raises(ValueError):
backend.load_pipeline("some/unrecognised-repo", gguf_filename = "x.gguf")
# load_progress state machine (no threads / network / real cache)
from core.inference.diffusion import _LoadingState, _LoadState # noqa: E402
def test_load_progress_idle_and_ready():
backend = DiffusionBackend()
assert backend.load_progress()["phase"] is None
backend._state = _LoadState(object(), None, "r", "b", "cpu", "float32", False)
assert backend.load_progress()["phase"] == "ready"
def test_load_progress_error():
backend = DiffusionBackend()
backend._loading = _LoadingState(repo_id = "r", base_repo = "b", error = "boom")
p = backend.load_progress()
assert p["phase"] == "error" and p["error"] == "boom"
def test_load_progress_downloading_then_finalizing(monkeypatch):
backend = DiffusionBackend()
backend._loading = _LoadingState(repo_id = "r", base_repo = "b", expected_bytes = 1000)
monkeypatch.setattr(DiffusionBackend, "_cache_bytes", staticmethod(lambda repo: 150))
p = backend.load_progress()
assert p["phase"] == "downloading"
assert p["bytes_downloaded"] == 300 # summed across repo + base
assert abs(p["fraction"] - 0.3) < 1e-9
monkeypatch.setattr(DiffusionBackend, "_cache_bytes", staticmethod(lambda repo: 500))
assert backend.load_progress()["phase"] == "finalizing" # 1000/1000
def test_base_file_downloaded_excludes_undownloaded():
# Counted: the pipeline manifest + component subfolders from_pretrained fetches.
assert _base_file_downloaded("model_index.json")
assert _base_file_downloaded("text_encoder/model-00001-of-00003.safetensors")
assert _base_file_downloaded("vae/diffusion_pytorch_model.safetensors")
# Excluded: the GGUF supplies the transformer; docs/assets and top-level files
# are never downloaded, so counting them would peg the bar short of 100%.
assert not _base_file_downloaded(
"transformer/diffusion_pytorch_model-00001-of-00003.safetensors"
)
assert not _base_file_downloaded("assets/Z-Image-Gallery.pdf")
assert not _base_file_downloaded("README.md")
assert not _base_file_downloaded(".gitattributes")
def test_load_progress_fraction_clamped(monkeypatch):
# The cache scan can exceed the estimate (e.g. a second cached quant); the
# reported fraction must still clamp to 1.0 rather than overshoot.
backend = DiffusionBackend()
backend._loading = _LoadingState(repo_id = "r", base_repo = "b", expected_bytes = 1000)
monkeypatch.setattr(DiffusionBackend, "_cache_bytes", staticmethod(lambda repo: 900))
p = backend.load_progress() # summed 1800 > expected 1000
assert p["phase"] == "finalizing"
assert p["fraction"] == 1.0
assert p["bytes_downloaded"] == 1000 # clamped to the estimate
def test_estimate_eta():
from core.inference.diffusion import _estimate_eta
# No rate yet until a step has elapsed since the first.
assert _estimate_eta(8, 1, first_step_at = 100.0, now = 100.0) is None
assert _estimate_eta(8, 0, first_step_at = 0.0, now = 100.0) is None
# 3 steps in 3s since the first ⇒ 1s/step ⇒ 4 steps left ⇒ ~4s.
assert _estimate_eta(8, 4, first_step_at = 100.0, now = 103.0) == 4.0
# Last step ⇒ 0 remaining.
assert _estimate_eta(8, 8, first_step_at = 100.0, now = 107.0) == 0.0
def test_generate_qwen_uses_true_cfg_scale(fake_runtime, tmp_path):
(tmp_path / "model.gguf").write_bytes(b"weights")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
base_repo = "Qwen/Qwen-Image",
family_override = "qwen-image",
)
backend.generate(prompt = "a sloth", guidance = 4.0)
# Qwen-Image's distilled guidance is off; the real CFG must land on true_cfg_scale.
call = backend._state.pipe.last_kwargs
assert call["true_cfg_scale"] == 4.0 and call["guidance_scale"] is None
def _load_ideogram(backend, tmp_path):
# Ideogram 4 loads only as a full pipeline (its two DiTs are assembled per-component
# by the stubbed load_ideogram4_pipeline); a local pipeline dir is enough here.
(tmp_path / "model_index.json").write_text("{}")
backend.load_pipeline(str(tmp_path), family_override = "ideogram-4")
def test_ideogram_rejects_single_file_and_gguf_kinds(fake_runtime, tmp_path):
# Ideogram 4 needs two DiTs assembled per-component, so there is no transformer-only
# single-file or GGUF load: the explicit kinds must be rejected up front (before a
# load evicts a working model), not assembled into a pipeline missing its second DiT.
backend = DiffusionBackend()
(tmp_path / "model.gguf").write_bytes(b"x")
with pytest.raises(ValueError, match = "full diffusers pipeline"):
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", family_override = "ideogram-4"
)
(tmp_path / "model.safetensors").write_bytes(b"x")
with pytest.raises(ValueError, match = "full diffusers pipeline"):
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.safetensors",
model_kind = "single_file",
family_override = "ideogram-4",
)
def test_generate_ideogram_defaults_keep_recommended_schedule(fake_runtime, tmp_path):
# Ideogram 4's pipeline defaults to its recommended tapered guidance_schedule
# (45x7.0 + 3x3.0, valid only at 48 steps) and REJECTS guidance_scale while the
# schedule is set. At the family's advertised defaults the backend must drop the
# constant so the recommended taper engages.
backend = DiffusionBackend()
_load_ideogram(backend, tmp_path)
backend.generate(prompt = "a sloth", steps = 48, guidance = 7.0)
call = backend._state.pipe.last_kwargs
assert call["guidance_scale"] is None # not passed: the pipe default engages
assert "guidance_schedule" not in call
def test_generate_ideogram_custom_guidance_nulls_schedule(fake_runtime, tmp_path):
# Any non-default request must broadcast the constant legally: guidance_scale set
# AND guidance_schedule explicitly nulled (the pipeline raises when both are set,
# and its default schedule is non-None).
backend = DiffusionBackend()
_load_ideogram(backend, tmp_path)
backend.generate(prompt = "a sloth", steps = 20, guidance = 5.0)
call = backend._state.pipe.last_kwargs
assert call["guidance_scale"] == 5.0
assert "guidance_schedule" in call and call["guidance_schedule"] is None
def test_begin_load_rejects_concurrent(monkeypatch):
backend = DiffusionBackend()
# The worker resolves the base + downloads, both over the network; stub them
# so the test is offline.
monkeypatch.setattr("core.inference.diffusion._hf_base_model", lambda *a, **k: None)
monkeypatch.setattr(DiffusionBackend, "_prefetch_files", lambda self, *a, **k: None)
monkeypatch.setattr(
DiffusionBackend, "_estimate_download_bytes", staticmethod(lambda *a, **k: (0, []))
)
# Block the spawned worker so the load stays "in progress".
monkeypatch.setattr(
DiffusionBackend, "load_pipeline", lambda self, **k: __import__("time").sleep(0.2)
)
backend.begin_load("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "z-image-turbo-Q4_K_S.gguf")
with pytest.raises(RuntimeError):
backend.begin_load("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "z-image-turbo-Q4_K_S.gguf")
def test_unload_cancels_in_flight_load(fake_runtime):
# An unload (or an arbiter eviction, which calls unload) while a load's worker
# is still resolving/downloading must cancel it: load_pipeline sees the bumped
# token and aborts, so the evicted load never resurrects a pipeline into VRAM.
backend = DiffusionBackend()
fam = detect_family("unsloth/Z-Image-Turbo-GGUF")
token = 7
backend._load_token = token
with pytest.raises(RuntimeError, match = "cancelled"):
# Simulate the worker reaching load_pipeline after unload bumped the token.
backend._load_token = token + 1
backend.load_pipeline(
"unsloth/Z-Image-Turbo-GGUF",
gguf_filename = "z-image-turbo-Q4_K_S.gguf",
base_repo = fam.base_repo,
_load_token = token,
)
def test_superseded_load_does_not_cancel_live_generation(fake_runtime):
# A superseded background load (its token was bumped by a newer load/unload) that
# finally reaches load_pipeline must bail WITHOUT signalling the current model's
# in-flight generation: the token check has to run before the cancel is set, or a
# stale worker aborts an unrelated, still-live denoise.
import threading as _threading
backend = DiffusionBackend()
fam = detect_family("unsloth/Z-Image-Turbo-GGUF")
live_cancel = _threading.Event()
backend._active_generate_cancel = live_cancel # a generation from the CURRENT model
token = 11
backend._load_token = token + 1 # this load has already been superseded
with pytest.raises(RuntimeError, match = "cancelled"):
backend.load_pipeline(
"unsloth/Z-Image-Turbo-GGUF",
gguf_filename = "z-image-turbo-Q4_K_S.gguf",
base_repo = fam.base_repo,
_load_token = token,
)
assert not live_cancel.is_set() # the live generation was left untouched
def test_pick_dtype_bf16_only_on_ampere(fake_runtime, monkeypatch):
# BF16 only on Ampere+ (cc >= 8); pre-Ampere cards must fall back to FP16.
torch = sys.modules["torch"]
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)
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: (7, 5), raising = False)
assert backend._pick_device_and_dtype() == ("cuda", torch.float16)
def test_unload_sets_cancel_event(fake_runtime):
# unload signals an in-flight download (which runs without the lock) to abort.
backend = DiffusionBackend()
assert not backend._cancel_event.is_set()
backend.unload()
assert backend._cancel_event.is_set()
def test_prefetch_aborts_when_cancelled(tmp_path):
# A prefetch interrupted by unload (cancel event set) raises rather than
# downloading the whole base, so the load can be preempted mid-download.
backend = DiffusionBackend()
backend._cancel_event.set()
# Local gguf path so the transformer download is skipped; the base loop hits
# the cancel check on its first file (no network).
(tmp_path / "model.gguf").write_bytes(b"x")
with pytest.raises(RuntimeError, match = "Cancelled"):
backend._prefetch_files(
str(tmp_path),
"model.gguf",
"Tongyi-MAI/Z-Image-Turbo",
["vae/diffusion_pytorch_model.safetensors"],
None,
)
def test_prefetch_downloads_gguf_and_base(monkeypatch, tmp_path):
backend = DiffusionBackend()
calls: list = []
monkeypatch.setattr(
"utils.hf_xet_fallback.hf_hub_download_with_xet_fallback",
lambda repo, fn, tok, **k: (calls.append((repo, fn)), f"/cache/{fn}")[1],
)
# Hub repo: the GGUF transformer and each base file are fetched.
backend._prefetch_files(
"unsloth/Z-Image-Turbo-GGUF",
"model.gguf",
"base/repo",
["vae/x.safetensors", "text_encoder/y.safetensors"],
"hf_tok",
)
assert ("unsloth/Z-Image-Turbo-GGUF", "model.gguf") in calls
assert ("base/repo", "vae/x.safetensors") in calls
assert ("base/repo", "text_encoder/y.safetensors") in calls
# Local GGUF path: the transformer download is skipped, base still fetched.
calls.clear()
(tmp_path / "model.gguf").write_bytes(b"x")
backend._prefetch_files(str(tmp_path), "model.gguf", "base/repo", ["vae/x.safetensors"], None)
assert all(repo != str(tmp_path) for repo, _ in calls)
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
def test_bad_mode_strings_fail_before_eviction(fake_runtime):
# Every mode normalizer that can raise runs BEFORE the load evicts the previous
# pipeline, so a bad request never costs the user their working model.
backend = DiffusionBackend()
fam = detect_family("unsloth/Z-Image-GGUF")
backend._state = _LoadState(
pipe = object(),
family = fam,
repo_id = "r",
base_repo = "b",
device = "cpu",
dtype = "float32",
cpu_offload = False,
)
for kwargs in (
{"transformer_quant": "int7"},
{"speed_mode": "warp"},
{"attention_backend": "bogus"},
{"transformer_cache": "bogus"},
{"text_encoder_quant": "fp3"},
):
with pytest.raises(ValueError):
backend.load_pipeline("unsloth/Z-Image-GGUF", gguf_filename = "m.gguf", **kwargs)
assert backend._state is not None
# Lock split + mid-denoise cancellation
def test_generate_lock_split_keeps_status_and_unload_responsive(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 only _generate_lock
# status() / generate_progress() must NOT block behind the denoise.
assert backend.status()["loaded"] is True
assert backend.generate_progress()["active"] is True
cancel_ref = backend._active_generate_cancel
assert cancel_ref is not None
# unload() signals THIS generation's cancel event, then waits for the denoise to
# actually exit before returning: callers treat its return as "VRAM is free" (the
# GPU arbiter hands the GPU to chat on it). Release the pipe once the cancel
# lands, standing in for the step callback of a real pipeline.
releaser = threading.Thread(target = lambda: (cancel_ref.wait(5), release.set()))
releaser.start()
backend.unload()
releaser.join(5)
assert cancel_ref.is_set()
assert backend.status()["loaded"] is False
t.join(5)
# The cancelled generation raised rather than returning a now-evicted image, and
# it had already exited (deregistering its cancel) before unload() returned.
assert "exc" in out and "cancelled" in str(out["exc"]).lower()
assert backend._active_generate_cancel is None
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()
# No filename + unsloth repo -> a full-pipeline load (allowed for unsloth/*).
assert backend.validate_load_request("unsloth/Z-Image-Turbo-unsloth-bnb-4bit").name == "z-image"
# No filename + non-unsloth repo -> a pipeline load, gated to unsloth/* -> rejected.
with pytest.raises(ValueError, match = "unsloth"):
backend.validate_load_request("some-org/Z-Image-bnb-4bit")
# An explicit gguf/single_file kind still requires a single-file name.
with pytest.raises(ValueError, match = "single-file"):
backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", model_kind = "gguf")
# A pipeline kind must NOT carry a single-file name.
with pytest.raises(ValueError, match = "pipeline"):
backend.validate_load_request(
"unsloth/Z-Image-Turbo-bnb-4bit", gguf_filename = "q.gguf", model_kind = "pipeline"
)
# A single-file safetensors load is also gated to unsloth/* repos.
with pytest.raises(ValueError, match = "unsloth"):
backend.validate_load_request("some-org/Z-Image", gguf_filename = "model.safetensors")
with pytest.raises(ValueError, match = "family"):
backend.validate_load_request("meta/Llama-3", gguf_filename = "q.gguf")
# A family-looking repo paired with a non-GGUF single-file name is rejected here,
# BEFORE the route evicts chat and hands over the GPU (the background load would
# otherwise be the first to notice README.md is not a checkpoint).
with pytest.raises(ValueError, match = r"\.gguf"):
backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "README.md")
assert (
backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", gguf_filename = "q.gguf").name
== "z-image"
)
# A kind/extension mismatch fails fast here, before the route evicts chat + grabs the
# GPU only to fail in the background from_single_file path.
with pytest.raises(ValueError, match = ".gguf"):
backend.validate_load_request(
"unsloth/Z-Image-Turbo-GGUF", gguf_filename = "model.safetensors", model_kind = "gguf"
)
with pytest.raises(ValueError, match = "gguf"):
backend.validate_load_request(
"unsloth/Qwen-Image-2512-FP8", gguf_filename = "q.gguf", model_kind = "single_file"
)
# A remote "*-GGUF" repo loaded as a full pipeline (no single-file name) is a single-file
# GGUF repo, so from_pretrained would find no pipeline manifest and fail after chat is
# already evicted; reject it here before the GPU handoff.
with pytest.raises(ValueError, match = "GGUF"):
backend.validate_load_request("unsloth/Z-Image-Turbo-GGUF", model_kind = "pipeline")
# 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",
)
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)
# ── Phase 2A: memory policy wiring (load -> planner -> placement) ──────────────
def test_load_reports_memory_plan_fields_on_cpu(fake_runtime, tmp_path):
# The default stub resolves to a CPU target: no offload is possible, but VAE
# tiling is on (no separate device pool), and status carries the new fields.
(tmp_path / "m.gguf").write_bytes(b"weights")
backend = DiffusionBackend()
status = backend.load_pipeline(str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image")
assert status["offload_policy"] == "none"
assert status["cpu_offload"] is False
assert status["vae_tiling"] is True
assert status["memory_mode"] == "auto"
pipe = backend._state.pipe
assert pipe.moved_to == "cpu" and pipe.vae_tiled and pipe.vae_sliced
def _force_cuda_target(backend, monkeypatch):
"""Drive the loader down the CUDA (offload-capable) path under the stub."""
torch = sys.modules["torch"]
monkeypatch.setattr(backend, "_pick_device_and_dtype", lambda: ("cuda", torch.bfloat16))
def test_load_memory_mode_balanced_streams_or_falls_back(fake_runtime, tmp_path, monkeypatch):
# balanced requests streamed block-level (group) offload. Under the stub there is
# no real diffusers.hooks, so group can't engage and the applier falls back to
# whole-module offload, reporting the policy actually engaged (the real "group"
# path is GPU-verified in the bench).
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", memory_mode = "balanced"
)
assert status["offload_policy"] in ("group", "model") and status["cpu_offload"] is True
assert status["memory_mode"] == "balanced"
assert backend._state.pipe.offloaded is True # model-offload fallback engaged
def test_load_memory_mode_low_vram_engages_model_offload(fake_runtime, tmp_path, monkeypatch):
# low_vram offloads every component (lowest VRAM); whole-module offload is the
# robust path and engages directly (no streaming, so no diffusers.hooks needed).
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", memory_mode = "low_vram"
)
assert status["offload_policy"] == "model" and status["cpu_offload"] is True
pipe = backend._state.pipe
assert pipe.offloaded is True and pipe.moved_to is None # offload owns placement
def test_load_explicit_cpu_offload_engages_model_offload_on_cuda(
fake_runtime, tmp_path, monkeypatch
):
# cpu_offload=True with no mode: auto would stay resident (budget unknown under
# the stub), but the explicit flag forces whole-module offload.
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", cpu_offload = True
)
assert status["offload_policy"] == "model" and status["cpu_offload"] is True
def test_load_speed_mode_gguf_auto_defaults_and_explicit(fake_runtime, tmp_path):
# No speed_mode on a GGUF model -> auto `default` (near-lossless, compile sits
# below the quant noise floor). compile itself only engages on CUDA, so on this
# CPU stub no optim need engage, but the resolved mode is `default`.
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image")
assert status["speed_mode"] == "default"
# An explicit "off" opts back into the bit-identical path (engages nothing).
status_off = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", speed_mode = "off"
)
assert status_off["speed_mode"] == "off" and status_off["speed_optims"] == []
# An explicit speed_mode threads through to status (engaged optims are GPU-verified).
status2 = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", speed_mode = "max"
)
assert status2["speed_mode"] == "max"
# Text-encoder quant defaults off (None); a requested mode threads through (the
# actual engagement is GPU-verified, since it needs real torch/torchao).
assert status2["text_encoder_quant"] is None
status3 = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
text_encoder_quant = "nvfp4",
)
# Under the CPU stub nvfp4 is unsupported, so it engages nothing -> None.
assert status3["text_encoder_quant"] is None
def test_load_fast_mode_stays_resident_on_cuda(fake_runtime, tmp_path, monkeypatch):
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image", memory_mode = "fast"
)
assert status["offload_policy"] == "none" and status["cpu_offload"] is False
assert backend._state.pipe.moved_to == "cuda"
# ── transformer quant (opt-in dense fast path) ────────────────────────────────
def _stub_dense_quant(monkeypatch, *, scheme = "fp8"):
"""Force the dense+quant branch hermetically: a supported dense source, a
from_pretrained on the fake transformer, and a quantizer that engages `scheme`.
Returns a dict recording the dense-loader / quantizer calls."""
from core.inference import diffusion as dmod
calls: dict = {"from_pretrained": 0, "quantize": 0, "quant_mode": None}
@classmethod
def _from_pretrained(cls, base, **kwargs):
calls["from_pretrained"] += 1
calls["fp_kwargs"] = {"base": base, **kwargs}
return object()
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _from_pretrained, raising = False)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
# Resolve the scheme without the real GPU smoke probe, and configure no pre-quant
# checkpoint so the dense materialise+quantise branch is the one exercised.
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: scheme
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
def _quantize(pipe, target, *, mode, **kw):
calls["quantize"] += 1
calls["quant_mode"] = mode
return scheme
monkeypatch.setattr(dmod, "quantize_transformer", _quantize)
return calls
def test_default_load_autos_dense_gate_and_falls_back(fake_runtime, tmp_path, monkeypatch):
# UNSET Dtype defaults to the hardware ladder: the dense gate IS consulted, and a
# device without dense support (this fake runtime) falls back to the GGUF build.
from core.inference import diffusion as dmod
consulted = {"n": 0}
def _supported(*a, **k):
consulted["n"] += 1
return False
monkeypatch.setattr(dmod, "dense_transformer_supported", _supported)
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image")
assert consulted["n"] >= 1
assert status["transformer_quant"] is None
assert _FakeTransformer.last["path"] # GGUF from_single_file was used
def test_explicit_off_load_skips_dense_quant_path(fake_runtime, tmp_path, monkeypatch):
# An EXPLICIT "none" pins running the GGUF as-is: the dense gate is never even
# consulted (short-circuit), so the pinned-off contract cannot regress.
from core.inference import diffusion as dmod
monkeypatch.setattr(
dmod,
"dense_transformer_supported",
lambda *a, **k: pytest.fail("dense path must not run with an explicit off"),
)
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "none",
)
assert status["transformer_quant"] is None
assert _FakeTransformer.last["path"] # GGUF from_single_file was used
def test_transformer_quant_dense_path_engaged(fake_runtime, tmp_path, monkeypatch):
# transformer_quant + a CUDA resident plan -> load the DENSE transformer from the
# base repo, place it on the device, quantise it, and report the engaged scheme.
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
calls = _stub_dense_quant(monkeypatch, scheme = "fp8")
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
)
assert status["transformer_quant"] == "fp8"
# No speed_mode was given, but a quantized transformer is ~30x slower eager, so the
# backend promotes it to `default` (regional compile) instead of the dense `off`.
assert status["speed_mode"] == "default"
assert calls["from_pretrained"] == 1 and calls["quantize"] == 1
assert calls["quant_mode"] == "fp8"
assert calls["fp_kwargs"]["subfolder"] == "transformer" # dense transformer subfolder
# The GGUF single-file path was NOT used for the transformer.
assert _FakeTransformer.last == {}
# quantize ran on-device: the dense pipe was placed on cuda (before compile).
assert backend._state.pipe.moved_to == "cuda"
assert status["offload_policy"] == "none"
def test_transformer_quant_prequant_path_engaged(fake_runtime, tmp_path, monkeypatch):
# A configured pre-quant checkpoint -> load the already-quantized transformer directly;
# the dense from_pretrained and the on-device quantize_transformer are NOT used.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
prequant_obj = object()
loaded: dict = {"n": 0}
def _load_prequant(transformer_cls, base, source, **kw):
loaded["n"] += 1
loaded["scheme"] = kw.get("scheme")
return prequant_obj
monkeypatch.setattr(dmod, "load_prequantized_transformer", _load_prequant)
@classmethod
def _fp_fail(cls, *a, **k):
pytest.fail("dense from_pretrained must not run when a prequant checkpoint loads")
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _fp_fail, raising = False)
monkeypatch.setattr(
dmod,
"quantize_transformer",
lambda *a, **k: pytest.fail("quantize_transformer must not run on the prequant path"),
)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
transformer_prequant_path = str(tmp_path / "zimage_fp8.pt"),
)
assert status["transformer_quant"] == "fp8"
assert loaded["n"] == 1 and loaded["scheme"] == "fp8"
# The pre-quantized transformer object was assembled into the pipeline...
assert _FakePipeline.last.get("transformer") is prequant_obj
# ...and the GGUF single-file path was not used.
assert _FakeTransformer.last == {}
def test_transformer_quant_prequant_load_fails_falls_back_to_dense(
fake_runtime, tmp_path, monkeypatch
):
# A configured prequant source whose load returns None must fall back to the dense
# materialise+quantise path (not straight to GGUF), preserving the fast mode.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
calls = _stub_dense_quant(monkeypatch, scheme = "fp8")
# Override the no-prequant default: a source resolves, but its load fails.
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
monkeypatch.setattr(dmod, "load_prequantized_transformer", lambda *a, **k: None)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
)
assert status["transformer_quant"] == "fp8"
assert calls["from_pretrained"] == 1 and calls["quantize"] == 1 # dense path ran
assert _FakeTransformer.last == {} # GGUF not used
def test_transformer_quant_falls_back_to_gguf_on_failure(fake_runtime, tmp_path, monkeypatch):
# A dense/quant failure (here: quantize returns None -> unsupported) must fall back
# to the GGUF build, not error -- status reports no transformer_quant engaged.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
@classmethod
def _from_pretrained(cls, base, **kwargs):
return object()
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _from_pretrained, raising = False)
monkeypatch.setattr(dmod, "quantize_transformer", lambda pipe, target, **kw: None)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
)
assert status["loaded"] is True
assert status["transformer_quant"] is None # fell back
assert _FakeTransformer.last["path"] # GGUF from_single_file used
def test_transformer_quant_skipped_when_plan_offloads(fake_runtime, tmp_path, monkeypatch):
# The dense bf16 transformer only fits resident, so when the memory plan would
# offload (here low_vram) the fast path is skipped and GGUF loads instead -- the
# dense transformer is never even loaded.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
@classmethod
def _fp_fail(cls, *a, **k):
pytest.fail("dense transformer must not load when the plan offloads")
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _fp_fail, raising = False)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
memory_mode = "low_vram",
)
assert status["transformer_quant"] is None
assert status["offload_policy"] == "model"
assert _FakeTransformer.last["path"] # GGUF path used
def test_transformer_quant_unsupported_scheme_skips_dense_download(
fake_runtime, tmp_path, monkeypatch
):
# An explicit unsupported scheme (select_transformer_quant_scheme -> None) must fail
# the dense path BEFORE materialising the multi-GB dense transformer, then fall back
# to GGUF -- otherwise the download runs under the load lock during finalization
# after the old model was already evicted, only to fail at quantize.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: None
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
@classmethod
def _fp_fail(cls, *a, **k):
pytest.fail("dense transformer must not download when the scheme is unsupported")
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _fp_fail, raising = False)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path),
gguf_filename = "m.gguf",
family_override = "z-image",
transformer_quant = "fp8",
)
assert status["loaded"] is True
assert status["transformer_quant"] is None # fell back to GGUF
assert _FakeTransformer.last["path"] # GGUF from_single_file used
def test_base_file_downloaded_include_transformer_flag():
# Default: transformer/ shards are the GGUF's job, so they are excluded from
# the prefetch list; the dense transformer-quant path opts them back in.
from core.inference.diffusion import _base_file_downloaded
assert _base_file_downloaded("transformer/diffusion_pytorch_model-00001.safetensors") is False
assert (
_base_file_downloaded(
"transformer/diffusion_pytorch_model-00001.safetensors", include_transformer = True
)
is True
)
# The flag must not admit anything else that is normally excluded.
assert _base_file_downloaded("assets/teaser.png", include_transformer = True) is False
assert _base_file_downloaded("README.md", include_transformer = True) is False
def test_dense_quant_prefetch_needed_gates(fake_runtime, monkeypatch):
# The transformer/ prefetch only widens when the dense quant path can really
# run: quant requested + device supported + scheme resolvable + no prequant
# checkpoint shortcutting the dense build.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
fam = detect_family("unsloth/Z-Image-Turbo-GGUF")
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is True
# UNSET defaults to the hardware ladder (Dtype default-auto) -> widens too.
assert backend._dense_quant_prefetch_needed(fam, {}) is True
# An explicit off pins running the GGUF as-is -> never widen.
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "none"}) is False
# A resolvable pre-quantized checkpoint shortcuts the dense download.
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False
# Unsupported scheme bails before the dense path (and so must the prefetch).
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: None
)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False
# Device without dense support (e.g. non-CUDA) never widens.
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: False)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False
def test_companion_cache_bytes_local_dir_excludes_transformer(tmp_path):
# A LOCAL diffusers base: sum the on-disk VAE / text-encoder weights so auto memory
# planning sees the resident companions, but exclude transformer/ (the GGUF supplies
# it) and non-weight files. A folded-to-zero companion could OOM a resident plan.
(tmp_path / "vae").mkdir()
(tmp_path / "vae" / "diffusion_pytorch_model.safetensors").write_bytes(b"x" * 100)
(tmp_path / "text_encoder").mkdir()
(tmp_path / "text_encoder" / "model.safetensors").write_bytes(b"y" * 50)
(tmp_path / "transformer").mkdir()
(tmp_path / "transformer" / "diffusion_pytorch_model.safetensors").write_bytes(b"z" * 9999)
(tmp_path / "model_index.json").write_bytes(b"{}") # non-weight file, ignored
total = DiffusionBackend._companion_cache_bytes(str(tmp_path))
assert total == 150 # vae + text_encoder only; transformer/ and json excluded
def test_reset_step_cache_helper_is_best_effort():
# Calls the transformer's reset hook when present.
calls = []
pipe = types.SimpleNamespace(
transformer = types.SimpleNamespace(reset_stateful_hooks = lambda: calls.append(True))
)
DiffusionBackend._reset_step_cache(pipe)
assert calls == [True]
# No transformer, or a transformer without the hook -> silent no-op (never raises).
DiffusionBackend._reset_step_cache(types.SimpleNamespace())
DiffusionBackend._reset_step_cache(types.SimpleNamespace(transformer = object()))
def test_generate_resets_step_cache_only_when_engaged(fake_runtime, tmp_path):
# FBCache residuals live on the resident transformer across generations, so each
# generate() must reset the stateful cache first -- but only when a cache is engaged.
(tmp_path / "model.gguf").write_bytes(b"weights")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path),
gguf_filename = "model.gguf",
base_repo = "base/repo",
family_override = "z-image",
)
resets = []
backend._state.pipe.transformer = types.SimpleNamespace(
reset_stateful_hooks = lambda: resets.append(True)
)
# No cache engaged (transformer_cache is None) -> reset must NOT run.
backend.generate(prompt = "a sloth")
assert resets == []
# Engage a cache; every subsequent generation resets the stateful cache first.
object.__setattr__(backend._state, "transformer_cache", "fbcache")
backend.generate(prompt = "a sloth")
backend.generate(prompt = "another sloth")
assert resets == [True, True]
def test_prefetch_returns_snapshot_dir_for_manifest(monkeypatch):
# The prefetched pipeline manifest's directory is the local snapshot root; a
# config-only base list (no manifest) returns None so the hub id stays in use.
backend = DiffusionBackend()
monkeypatch.setattr(
"utils.hf_xet_fallback.hf_hub_download_with_xet_fallback",
lambda repo, fn, tok, **k: f"/cache/snap/{fn}",
)
root = backend._prefetch_files(
"base/repo", None, "base/repo", ["model_index.json", "vae/x.safetensors"], None
)
assert root == "/cache/snap"
assert (
backend._prefetch_files("base/repo", None, "base/repo", ["vae/x.safetensors"], None) is None
)
def test_pipeline_load_uses_predownloaded_dir(fake_runtime, tmp_path):
# With a prefetched snapshot, from_pretrained must receive the local dir --
# its own hub sweep would re-download the root packaged singles the scoped
# prefetch skips (24 GB per FLUX.1 repo).
backend = DiffusionBackend()
backend.load_pipeline(
"unsloth/Qwen-Image-2512-bnb-4bit",
model_kind = "pipeline",
_base_local_dir = str(tmp_path),
)
assert _FakePipeline.last["base"] == str(tmp_path)
backend.unload()