unsloth/studio/backend/tests/test_diffusion_backend.py
2026-07-05 04:42:18 +00:00

2154 lines
90 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"
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"):
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
# 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.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 _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 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_skips_dense_quant_path(fake_runtime, tmp_path, monkeypatch):
# With no transformer_quant flag the GGUF path is taken and the dense gate is
# never even consulted (short-circuit), so the default 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 without the flag"),
)
(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["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
# No quant requested -> never widen.
assert backend._dense_quant_prefetch_needed(fam, {}) 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()