unsloth/studio/backend/core/inference/diffusion.py
Daniel Han 6e16ad16f7 Tighten diffusion comments
Collapse the multi-line comment blocks across the image, video, sd.cpp and diffusion-training code to one or two lines each, and drop comments that only restate the statement below them. Comments only, no code or behaviour changes.
2026-07-27 11:51:20 +00:00

3457 lines
180 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
"""Local diffusion (text-to-image) backend.
A torch-only singleton that loads one of three "kinds" (see ``resolve_model_kind``):
a single-file GGUF transformer dequantised on-device via ``GGUFQuantizationConfig``,
a single-file safetensors transformer (e.g. fp8), or a full diffusers pipeline via
``from_pretrained`` (which re-applies an embedded quant config such as bnb-4bit). The
single-file kinds pull the rest of the pipeline (VAE, text encoders, scheduler) from
the matching base repo; the pipeline kind pulls everything from the repo itself.
Non-GGUF kinds are gated to the ``unsloth/*`` org (or a local path) for safety.
torch/diffusers are imported lazily so this stays importable in a no-torch runtime.
``begin_load`` runs on a background thread; poll ``load_progress`` for the download
bar. GPU-handoff policy lives in the arbiter the routes call, not here.
"""
from __future__ import annotations
import inspect
import json
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, Optional
from loggers import get_logger
from utils.hardware import clear_gpu_cache
from .diffusion_families import (
DIFFUSION_CANCELLED_MSG,
DIFFUSION_NOT_LOADED_MSG,
IDEOGRAM4_FAMILY_NAME,
LUMINA2_FAMILY_NAME,
DiffusionFamily,
assert_pipeline_class_available,
default_generation_params,
detect_family_for_pick,
excluded_model_reason,
resolve_base_repo,
resolve_local_gguf_child,
supported_family_names,
)
from .diffusion_device import (
DiffusionDeviceTarget,
diffusion_device_target_from_torch_device,
resolve_diffusion_device_target,
)
from .diffusion_ideogram4 import ideogram4_repo_is_fp8, load_ideogram4_pipeline
from .diffusion_hidream import HIDREAM_FAMILY_NAME, hidream_te4_kwargs
from .diffusion_krea2 import KREA2_FAMILY_NAME, load_krea2_pipeline
from .diffusion_memory import (
MEMORY_MODE_BALANCED,
MEMORY_MODE_LOW_VRAM,
OFFLOAD_NONE,
apply_memory_plan,
estimate_gguf_resident_mib,
estimate_image_runtime_mib,
estimate_safetensors_dense_mib,
file_size_mib,
normalize_memory_mode,
plan_diffusion_memory,
plan_fits_total_capacity,
settled_snapshot_device_memory,
)
from .diffusion_speed import (
SPEED_DEFAULT,
SPEED_MAX,
SPEED_OFF,
apply_speed_optims,
compile_eligible,
compiled_shapes_are_static,
normalize_speed_mode,
resolve_speed_mode,
restore_backend_flags,
snapshot_backend_flags,
)
from .diffusion_attention import (
apply_attention_backend,
normalize_attention_backend,
select_attention_backend,
_ensure_attention_backend_installed,
)
from . import diffusion_compile_cache as compile_cache
from . import diffusion_cond_cache as cond_cache
from . import diffusion_gguf_compile as gguf_compile
from .diffusion_batched import (
chunk_jobs,
is_oom_error,
resolve_batch_jobs,
split_chunk,
uniform_prompt,
)
from .diffusion_cache import (
FBCACHE_MIN_STEPS,
TC_AUTO,
TC_FBCACHE,
apply_step_cache,
effective_denoise_steps,
effective_request_strength,
maybe_toggle_step_cache,
normalize_transformer_cache,
)
from .diffusion_precision import normalize_te_quant, quantize_text_encoders
from .diffusion_te_prequant import te_prequant_pipe_kwargs
from .diffusion_prequant import (
load_prequantized_transformer,
resolve_prequant_source,
usable_prequant_source,
)
from .diffusion_auto_policy import (
build_resolved_record,
family_bf16_components_gb,
resolve_dense_quant_candidate,
)
from .diffusion_transformer_quant import (
TQ_AUTO,
DEFAULT_MIN_LINEAR_FEATURES,
dense_transformer_supported,
normalize_transformer_quant,
quantize_transformer,
select_transformer_quant_scheme,
)
logger = get_logger(__name__)
# Load kinds: "gguf" (GGUF transformer dequantised on-device) and "single_file" (safetensors, e.g. fp8)
# both take companions from the base repo; "pipeline" is a full diffusers repo via from_pretrained.
_MODEL_KINDS = frozenset({"gguf", "single_file", "pipeline"})
def hub_cache_dir() -> str:
"""The cache root every loader call must be pinned to.
diffusers resolves an unset cache_dir through huggingface_hub's import-time constant,
which a mid-session cache-folder change does not update. The prefetch reads the live
setting, so without this a single load could split across two roots."""
from utils.hf_cache_settings import active_hf_hub_cache
return active_hf_hub_cache()
def resolve_model_kind(gguf_filename: Optional[str], model_kind: Optional[str] = None) -> str:
"""Classify a load request into one of ``_MODEL_KINDS``.
An explicit ``model_kind`` wins (validated). Otherwise the kind is inferred from
the single-file name: a ``.gguf`` name is ``"gguf"``, any other single-file name is
``"single_file"``, and the absence of a name is a full ``"pipeline"`` load. Pure and
network-free, so the route, validation, and load paths all agree on the kind."""
if model_kind:
kind = model_kind.strip().lower()
if kind not in _MODEL_KINDS:
raise ValueError(
f"Unknown model_kind '{model_kind}'. Expected one of {sorted(_MODEL_KINDS)}."
)
return kind
name = (gguf_filename or "").strip()
if not name:
return "pipeline"
if name.lower().endswith(".gguf"):
return "gguf"
return "single_file"
def _active_lora_pairs(pipe: Any) -> list:
"""``[(name, weight)]`` for the adapters actually attached to ``pipe``, zero-weight ones
dropped.
Reads the ``_unsloth_loras`` marker, which the LoRA paths write as ``(name, path, weight)``.
Shape is tolerated rather than assumed: this runs inside the generate result and an unpacking
error here would sink a finished generation whose images are already in hand."""
pairs = []
for entry in getattr(pipe, "_unsloth_loras", ()) or ():
try:
if len(entry) == 3:
name, _path, weight = entry
elif len(entry) == 2:
name, weight = entry
else:
continue
weight = float(weight)
except Exception: # noqa: BLE001 — an unrecognised marker records no adapter
continue
if weight:
pairs.append((name, weight))
return pairs
def resolve_local_single_file(model_path: str) -> Optional[str]:
"""The sole single-file checkpoint basename in a local ``model_path`` directory that is NOT a
diffusers pipeline (no ``model_index.json``) and holds exactly one ``.safetensors`` file, else
None.
The On-Device scanner advertises a bare single-file safetensors directory as a text-to-image
model (it matches a known family by name), but the local picker starts it as a ``pipeline``
with no filename, so a pipeline load 400s on the missing ``model_index.json`` and the
advertised model is unusable. The images load route uses this to reinterpret such a pick as a
``single_file`` load of the sole checkpoint. A real pipeline dir (has ``model_index.json``) or
an ambiguous one (0 or more than 1 ``.safetensors``, e.g. a sharded pipeline) returns None and
loads unchanged. A PEFT LoRA adapter folder is also skipped (see below). Never raises."""
try:
root = Path(model_path).expanduser()
if not root.is_dir() or (root / "model_index.json").is_file():
return None
# A PEFT adapter folder is not a base checkpoint; skip it so validation 400s before eviction.
if (root / "adapter_config.json").is_file():
return None
checkpoints = [
p.name
for p in root.iterdir()
if p.is_file()
and p.suffix.lower() == ".safetensors"
and p.stem.lower() != "adapter_model"
]
except OSError:
return None
return checkpoints[0] if len(checkpoints) == 1 else None
def _decode_b64_image(data: str, *, mode: str = "RGB") -> Any:
"""Decode a base64 (optionally ``data:`` URL) image string to a PIL image.
The image-conditioned workflows (img2img / inpaint / edit) transport the input
image and mask as base64 in the JSON request, so this is the single decode path.
A mask is decoded as single-channel ``L``; the source image as ``RGB``."""
import base64
import binascii
import io
from PIL import Image
raw = data.strip()
if raw.startswith("data:"):
# data:[<mime>][;base64],<payload>
_, _, raw = raw.partition(",")
try:
blob = base64.b64decode(raw, validate = False)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"Invalid base64 image data: {exc}") from exc
# Bound the decoded size: 4096px covers txt2img 2048, upscales and outpaint canvases.
max_side = 4096
try:
img = Image.open(io.BytesIO(blob))
# Reject from the header before img.load() so a huge-dimension file cannot spike memory.
w, h = img.size
if w > max_side or h > max_side:
raise ValueError(f"Image is too large ({w}x{h}); maximum is {max_side}px per side.")
img.load()
except ValueError:
raise # the size guard's own message; don't wrap it as a decode error
except Exception as exc: # noqa: BLE001 — surfaced as a 400 to the client
raise ValueError(f"Could not decode image: {exc}") from exc
return img.convert(mode)
def _snap_to_multiple(img: Any, multiple: int = 16) -> Any:
"""Resize a PIL image so both sides are multiples of ``multiple`` (rounded to nearest,
minimum one multiple), preserving content with a high-quality resample.
Image-conditioned pipelines (Z-Image / Qwen / FLUX: 8x VAE downsample + 2x patch) reject
sizes that are not divisible by 16. Rather than error on an odd-sized upload, snap it so
the workflow just works; rounding to nearest keeps the rescale minimal/accurate."""
from PIL import Image
w, h = img.size
nw = max(multiple, int(round(w / multiple)) * multiple)
nh = max(multiple, int(round(h / multiple)) * multiple)
if (nw, nh) != (w, h):
img = img.resize((nw, nh), Image.LANCZOS)
return img
def _clamp_max_side(img: Any, max_side: int) -> Any:
"""Downscale a PIL image so its longest side is <= ``max_side``, preserving aspect ratio
(high-quality resample); a no-op when it already fits.
img2img / inpaint take their OUTPUT size from the uploaded image, so without a bound an
oversized upload (up to the 4096/side decode cap -- 4x the txt2img 2048 ceiling, ~16x the
area) drives a proportionally larger latent and O(n^2) attention that OOMs the transformer/
VAE on a normal card, surfacing only as an opaque 500. Clamping the longest side to the same
2048 ceiling txt2img enforces (and upscale caps to) keeps these workflows bounded."""
from PIL import Image
w, h = img.size
longest = max(w, h)
if longest <= max_side:
return img
scale = max_side / float(longest)
nw = max(1, int(round(w * scale)))
nh = max(1, int(round(h * scale)))
return img.resize((nw, nh), Image.LANCZOS)
def _compile_shape_dims(workflow: str, init_pil: Any, width: int, height: int) -> tuple[int, int]:
"""The (width, height) a generation's forward ACTUALLY runs at, for static
compile-cache shape registration.
txt2img / reference / controlnet generate at the requested slider size, but the
image-conditioned workflows (img2img / inpaint / upscale / edit) derive the output
from the (resized/snapped) input image -- registering the slider values there would
mark a shape covered that was never compiled, so the truly-used shape never
re-dirties the bundle and warm restarts keep paying its compile. Mirrors the
width/height kwarg derivation in generate()."""
if workflow in ("txt2img", "reference", "controlnet") or init_pil is None:
return int(width), int(height)
iw, ih = init_pil.size
return int(iw), int(ih)
# Official non-unsloth bases loadable as a full pipeline: safetensors-only, no remote code, exact lowercased match (the img2img-only SDXL refiner is deliberately absent).
_TRUSTED_NON_GGUF_REPOS = frozenset(
{
"stabilityai/stable-diffusion-xl-base-1.0",
"stabilityai/sdxl-turbo",
# Vendor safetensors-only bases: LoRA training bases + the BF16 artifact per catalog group.
# FLUX.1 repos are Hub-gated (need the user's token); Qwen/Z-Image are open.
"black-forest-labs/flux.1-dev",
"black-forest-labs/flux.1-schnell",
"black-forest-labs/flux.1-kontext-dev",
# Krea: guidance-distilled FLUX.1-dev finetune, same arch and gating; detected via "flux.1".
"black-forest-labs/flux.1-krea-dev",
# FLUX.2 LoRA training bases (dev gated, klein-4B open); trusted because "Deploy to Create" reloads them as a pipeline.
"black-forest-labs/flux.2-dev",
"black-forest-labs/flux.2-klein-4b",
"tongyi-mai/z-image-turbo",
"qwen/qwen-image",
"qwen/qwen-image-2512",
"qwen/qwen-image-edit-2511",
# Krea 2: assembled per-component. Turbo = inference; Raw = the LoRA training base.
"krea/krea-2-turbo",
"krea/krea-2-raw",
# Lumina Image 2.0: standard diffusers layout, generic from_pretrained path.
"alpha-vllm/lumina-image-2.0",
# HunyuanImage 2.1: open community diffusers mirror, safetensors-only.
"hunyuanvideo-community/hunyuanimage-2.1-diffusers",
# HiDream-I1: open MIT weights, all three variants one family; Llama TE from the unsloth mirror.
"hidream-ai/hidream-i1-full",
"hidream-ai/hidream-i1-dev",
"hidream-ai/hidream-i1-fast",
# Ideogram 4: no bf16 ships. -fp8 stores both DiTs as raw float8 (the family base); both nf4 repos are identical bnb-4bit exports.
"ideogram-ai/ideogram-4-fp8",
"ideogram-ai/ideogram-4-nf4",
"ideogram-ai/ideogram-4-nf4-diffusers",
}
)
def _is_trusted_diffusion_repo(repo_id: str) -> bool:
"""Whether a NON-GGUF load is allowed for ``repo_id``.
Making ``gguf_filename`` optional opens a ``from_pretrained`` / ``from_single_file``
on an arbitrary repo, which fetches and deserialises third-party weights. So the
non-GGUF paths are gated to the ``unsloth/*`` org (the curated safetensors models),
a short allowlist of official safetensors-only base repos (``_TRUSTED_NON_GGUF_REPOS``,
e.g. the SDXL base), and local paths the user explicitly pointed at (already on their
disk). The GGUF path is unchanged and stays open to any repo, as before.
A bare ``owner/name`` HF id is never a real filesystem path, and an id with invalid
characters makes ``Path.exists()`` raise OSError; treat any such failure as "not a
local path" so the trust decision falls through to the org/allowlist checks (the
loader's validate_load_request raises the clear FileNotFoundError for a genuinely
missing local pick)."""
try:
if Path(repo_id).expanduser().exists():
return True
except OSError:
pass
rid = repo_id.strip().lower()
return rid.startswith("unsloth/") or rid in _TRUSTED_NON_GGUF_REPOS
def _assert_local_base_is_pipeline(base_repo: str) -> None:
"""A companion ``base_repo`` fed to ``from_pretrained(base)`` (or ``config=base``) must be a
diffusers PIPELINE directory (has ``model_index.json``). ``_is_trusted_diffusion_repo`` accepts
ANY existing local path, so without this a local base that is not a pipeline dir would pass the
preflight, let the route evict the resident GPU model, then fail deep in the background load --
the eviction this validation exists to prevent. A non-existent local base is already rejected
by the trust check (it is neither an existing path nor an unsloth/*/allowlisted repo); a bare
remote id is left for the loader to resolve. Shared by the image, video, and training preflights
so their local-base shape check stays in sync. Never evicts; raises ValueError on a bad local
base."""
base = (base_repo or "").strip()
if not base:
return
try:
root = Path(base).expanduser()
exists = root.exists()
except OSError:
return # invalid path characters -> a remote id, not a local path
if not exists:
return
if not root.is_dir() or not (root / "model_index.json").is_file():
raise ValueError(
f"Local base_repo is not a diffusers pipeline directory (no model_index.json): {base}"
)
@dataclass(frozen = True)
class _LoadState:
"""Everything about the currently-loaded pipeline, swapped as one unit."""
pipe: Any
family: Any
repo_id: str
base_repo: str
device: str
dtype: str
cpu_offload: bool
# Defaulted so older positional constructions keep working.
offload_policy: str = OFFLOAD_NONE
vae_tiling: bool = False
memory_mode: str = "auto"
# Resolved load kind ("gguf"|"single_file"|"pipeline"); lets the UI gate GGUF-only controls.
kind: str = "gguf"
speed_mode: str = SPEED_OFF
speed_optims: tuple = ()
# Torch backend flags (TF32 / cudnn.benchmark) captured before the speed layer mutated them.
backend_flags_before: Optional[dict] = None
# Text-encoder quant engaged: "fp8" | "nvfp4" | None.
text_encoder_quant: Optional[str] = None
# Transformer quant on the dense fast path ("int8"|"fp8"|"nvfp4"|"mxfp8"), or None when GGUF loaded.
transformer_quant: Optional[str] = None
# Attention backend via the diffusers dispatcher, or None for default SDPA.
attention_backend: Optional[str] = None
# Caller original attention request, so deferred engagement re-runs the same selection.
attention_request: Optional[str] = None
# Step cache engaged ("fbcache") or None. Opt-in, for many-step models.
transformer_cache: Optional[str] = None
# AUTO: generate() toggles FBCache across FBCACHE_MIN_STEPS; an explicit request never toggles.
cache_auto: bool = False
# Inputs the generation-time toggle re-applies (quantised threshold + override).
cache_quant_active: bool = False
cache_threshold: Optional[float] = None
# Shared eager patches installed for this load; uninstalled on unload.
eager_patched: bool = False
# Deferred speed auto: the load stays eager, generate() engages `default` at the 3rd generation.
speed_deferred: bool = False
# Successful generations on this load; drives the deferred engagement above.
generation_count: int = 0
# Pre-warmed torch.compile cache context when a compiled tier ran, else None.
compile_cache_ctx: Any = None
# Token kept so LoRA adapters selected at generate time can be fetched.
hf_token: Optional[str] = None
# Per-control provenance {control: {value, source, reason}}, for status badges.
resolved: Optional[dict] = None
@dataclass
class _LoadingState:
"""An in-flight background load, polled for download progress."""
repo_id: str
base_repo: str
expected_bytes: int = 0
error: Optional[str] = None
@dataclass
class _GenState:
"""An in-flight generation, updated per denoising step for the progress bar."""
total_steps: int
step: int = 0
# Set when the first step finishes, so the slower warmup step does not skew the ETA rate.
first_step_at: float = 0.0
# Computed once per step (in the callback) so it's stable between polls.
eta_seconds: Optional[float] = None
def _estimate_eta(total_steps: int, step: int, first_step_at: float, now: float) -> Optional[float]:
"""Seconds remaining, from the average step time measured after the first step.
None until at least one step has elapsed since the first."""
steps_since_first = step - 1
if not first_step_at or steps_since_first <= 0:
return None
per_step = (now - first_step_at) / steps_since_first
return max(0.0, (total_steps - step) * per_step)
def _resolve_diffusion_compute_dtype(fam: Optional[DiffusionFamily], dtype: Any) -> Any:
"""Promote float16 -> float32 for fp16-incompatible families (e.g. Z-Image),
whose activations overflow float16's finite range and render a black image.
Every other dtype/family passes through unchanged."""
if fam is None or not getattr(fam, "fp16_incompatible", False):
return dtype
import torch
return torch.float32 if dtype == torch.float16 else dtype
def _install_gguf_prefix_strip(transformer_cls: Any, logger: Any) -> None:
"""Wrap the class's diffusers single-file converter to strip the
``model.diffusion_model.`` container prefix that sd.cpp-converted GGUFs
carry on every tensor.
diffusers (<= 0.39) handles the prefix inconsistently: the FLUX.1 converter
strips it natively, but the FLUX.2 converter never does and KeyErrors on the
prefixed keys (``'double_blocks.0.img_attn.norm.key_norm'`` misparse in
``convert_flux2_transformer_checkpoint_to_diffusers``), and the Qwen-Image
mapping fn is an identity, so every prefixed tensor is reported "not used",
the model stays on meta, and ``.to(cuda)`` dies with "Cannot copy out of
meta tensor". Stripping the prefix when present is a no-op for already-clean
checkpoints, so the shim applies to every GGUF transformer class uniformly.
Idempotent (the wrapper is marked) and best-effort."""
try:
from diffusers.loaders import single_file_model as sfm
entry = sfm.SINGLE_FILE_LOADABLE_CLASSES.get(
getattr(transformer_cls, "__name__", str(transformer_cls))
)
if not isinstance(entry, dict):
return
original = entry.get("checkpoint_mapping_fn")
if not callable(original) or getattr(original, "_unsloth_prefix_strip", False):
return
prefix = "model.diffusion_model."
def _stripped_mapping_fn(checkpoint = None, **kwargs):
checkpoint = {
(key[len(prefix) :] if key.startswith(prefix) else key): value
for key, value in (checkpoint or {}).items()
}
return original(checkpoint = checkpoint, **kwargs)
_stripped_mapping_fn._unsloth_prefix_strip = True
entry["checkpoint_mapping_fn"] = _stripped_mapping_fn
except Exception as exc: # noqa: BLE001 — loader-compat shim only, never fail the load
logger.warning("diffusion.gguf: prefix-strip shim not installed: %s", exc)
class DiffusionBackend:
"""Holds at most one loaded diffusers pipeline. All mutations are serialised."""
def __init__(self) -> None:
# _lock serialises the small state mutations; the status/progress readers stay lock-free.
self._lock = threading.Lock()
# _generate_lock serialises generations and is the ONLY lock the denoise holds.
self._generate_lock = threading.Lock()
self._state: Optional[_LoadState] = None
self._loading: Optional[_LoadingState] = None
# Bumped on begin_load/unload so a superseded worker neither commits nor stamps progress.
self._load_token = 0
# Set by unload() to abort an in-flight (lock-free) download so an eviction preempts it.
self._cancel_event = threading.Event()
# Cancel Event of the in-flight generation; per-generation so a cancel can't be lost or leak.
self._active_generate_cancel: Optional[threading.Event] = None
# How many unloads / superseding loads are waiting on _generate_lock to free this pipeline.
# A generation queued behind the active one holds no cancel event yet, so the cancel they
# signal cannot reach it: without this fence it could win the lock as the active denoise
# released it, see a still-loaded _state, and run a whole new denoise after the model was
# told to go away -- stalling a chat/video GPU handoff for minutes and painting an image
# after an eject. A count, not a flag, so concurrent teardowns each own their own release.
self._teardown_waiters = 0
# Written by the callback, read lock-free by generate_progress().
self._gen: Optional[_GenState] = None
# img2img/inpaint pipes built via from_pipe (shared modules, no extra VRAM); cleared on unload.
self._aux_pipes: dict[str, Any] = {}
# Loaded ControlNets and their from_pipe pipelines, reusing resident modules; cleared on unload.
self._cn_models: dict[str, Any] = {}
self._cn_pipes: dict[tuple[str, str], Any] = {}
@property
def is_loaded(self) -> bool:
return self._state is not None
def _pick_device_and_dtype(self) -> tuple[str, Any]:
"""(device, dtype) for the current host. Thin wrapper over the device
policy module, kept as a method so tests can still monkeypatch it."""
target = resolve_diffusion_device_target()
return target.device, target.dtype
def _resolve_device_target(self, fam: Optional[DiffusionFamily]) -> DiffusionDeviceTarget:
"""The device target with the family fp16 guard applied.
Routes through _pick_device_and_dtype() (so a monkeypatched override still
drives the result), then promotes float16 -> float32 for fp16-incompatible
families (Z-Image), rebuilding the target so dtype + capability flags stay
consistent with the effective dtype.
"""
device, dtype = self._pick_device_and_dtype()
effective = _resolve_diffusion_compute_dtype(fam, dtype)
if effective is not dtype:
logger.warning(
"diffusion.dtype_promoted: family=%s float16 -> float32 (fp16-incompatible)",
getattr(fam, "name", None),
)
return diffusion_device_target_from_torch_device(device, effective)
def _resolve_gguf_path(self, repo_id: str, gguf_filename: str, hf_token: Optional[str]) -> str:
local_root = Path(repo_id).expanduser()
if local_root.exists():
return str(resolve_local_gguf_child(local_root, gguf_filename))
from huggingface_hub import hf_hub_download
return hf_hub_download(repo_id, gguf_filename, token = hf_token)
def _dense_quant_prefetch_needed(self, fam: DiffusionFamily, kwargs: dict) -> bool:
"""True when ``load_pipeline`` may take the dense transformer-quant path, so
the prefetch should also pull the base repo's ``transformer/`` shards.
Those shards are excluded from the prefetch by default (the GGUF supplies
the transformer), but ``_load_dense_quant_pipeline`` fetches them with
``from_pretrained(subfolder = "transformer")`` under the load lock during
"finalizing", after the previous pipeline was already evicted, where
unload/cancellation cannot preempt the download. Mirrors the dense-path
gates in ``load_pipeline``: quant requested and supported for this device,
and no pre-quantized checkpoint that would shortcut the dense build."""
raw = kwargs.get("transformer_quant")
# Unset defaults to the hardware ladder (mirrors load_pipeline's tri-state).
if raw is None or str(raw).strip().lower() in ("", "auto"):
mode = TQ_AUTO
else:
mode = normalize_transformer_quant(raw)
if mode is None:
return False
# An explicit Speed="off" load stays GGUF-as-is (dense path never runs); don't widen the prefetch.
speed = kwargs.get("speed_mode")
if speed is not None and str(speed).strip().lower() == SPEED_OFF:
return False
try:
# A definite-offload policy skips the dense build, so widening wastes a multi-GB pull with no GGUF fallback.
mm = normalize_memory_mode(kwargs.get("memory_mode"))
if mm in (MEMORY_MODE_BALANCED, MEMORY_MODE_LOW_VRAM):
return False
if mm is None and kwargs.get("cpu_offload"):
return False
target = self._resolve_device_target(fam)
# Only widen when the loader would take the dense path: resolve the same candidate load_pipeline re-plans against.
candidate = resolve_dense_quant_candidate(
fam = fam,
target = target,
requested = mode,
base_repo = kwargs.get("base_repo"),
prequant_path = kwargs.get("transformer_prequant_path"),
force_dense = bool(kwargs.get("loras")),
logger = None,
)
# A prequant loads a small checkpoint, so widening defeats the savings and can disk-full.
if candidate is None or candidate.prequant:
return False
# Capacity gate: mirror plan_fits_total_capacity against TOTAL capacity (not instantaneous free), since load_pipeline would otherwise decline the dense path anyway.
from .diffusion_memory import (
_reserve_mib,
snapshot_device_memory,
)
memory = snapshot_device_memory(target)
total = memory.total_mib
steady = getattr(candidate, "steady_total_mib", None)
if total is not None and steady is not None:
budget = int((int(total) - _reserve_mib(memory.memory_kind, int(total))) * 0.85)
if int(steady) > budget:
return False
return True
except Exception: # noqa: BLE001 — widening the prefetch is best-effort only
return False
def _prefetch_files(
self,
repo_id: str,
gguf_filename: Optional[str],
base: str,
base_files: list[str],
hf_token: Optional[str],
) -> Optional[str]:
"""Pre-download the GGUF + the given ``base_files`` into the HF cache,
WITHOUT the lock and honoring ``_cancel_event``, so load_pipeline's
from_single_file / from_pretrained hit the cache and the heavy download can
be preempted by an unload/eviction. Raises ``RuntimeError("Cancelled")``.
Returns the base repo's local snapshot dir when the prefetched set includes
the pipeline manifest, so from_pretrained can load from disk instead of
re-sweeping the hub (its own sweep also pulls files the scoped list skips,
e.g. the 24 GB packaged root singles in each FLUX.1 repo); None otherwise
(estimate failure, config-only base, local repo) -> hub id as before."""
from utils.hf_xet_fallback import hf_hub_download_with_xet_fallback
# GGUF transformer (hub repos only; a local path is already on disk).
if gguf_filename and not Path(repo_id).expanduser().exists():
hf_hub_download_with_xet_fallback(
repo_id, gguf_filename, hf_token, cancel_event = self._cancel_event
)
# Base repo (VAE / text-encoder / scheduler); list comes from the estimate.
snapshot_root: Optional[str] = None
for rfilename in base_files:
if self._cancel_event.is_set():
raise RuntimeError("Cancelled")
local = hf_hub_download_with_xet_fallback(
base, rfilename, hf_token, cancel_event = self._cancel_event
)
if rfilename == "model_index.json":
snapshot_root = str(Path(local).parent)
return snapshot_root
def validate_load_request(
self,
repo_id: str,
*,
gguf_filename: Optional[str] = None,
family_override: Optional[str] = None,
model_kind: Optional[str] = None,
base_repo: Optional[str] = None,
) -> DiffusionFamily:
"""Cheap, network-free validation shared by the route (before it evicts the
chat model) and the load paths, so an unloadable pick fails BEFORE the GPU
handoff. Resolves the load kind (gguf / single_file / pipeline), then raises
ValueError for a missing single-file name, a non-unsloth non-GGUF repo, or an
undetectable family, and ValueError/FileNotFoundError for a bad local path.
Touches no GPU, network, or state."""
kind = resolve_model_kind(gguf_filename, model_kind)
fam = detect_family_for_pick(repo_id, gguf_filename, family_override)
if fam is None:
# An excluded model gets its stated reason, not the unknown-family message that invites a doomed family_override retry.
excluded = excluded_model_reason(repo_id)
if excluded:
raise ValueError(f"'{repo_id}' cannot be loaded: {excluded}")
raise ValueError(
f"'{repo_id}' is not a supported diffusion image model. Supported families: "
f"{', '.join(supported_family_names())}. If this is a variant of one of them, "
f"pass family_override with that family name. (Video models and image models "
f"whose diffusers transformer has no single-file loader are not supported.)"
)
# Refuse a too-old diffusers here, not deep in the load after the checkpoint downloaded.
assert_pipeline_class_available(fam.pipeline_class, fam.name)
# Families whose single file IS the whole pipeline have no GGUF path; reject before eviction.
if kind == "gguf" and fam.single_file_is_pipeline:
raise ValueError(
f"'{fam.name}' checkpoints are whole-pipeline single files and have no GGUF "
f"transformer variant; load the .safetensors pipeline instead of a GGUF."
)
# A multi-denoiser family (Ideogram 4) has no transformer-only path; reject before eviction.
if kind in ("gguf", "single_file") and fam.pipeline_only:
raise ValueError(
f"'{fam.name}' loads only as a full diffusers pipeline (it assembles "
f"multiple transformers), not from a single-file or GGUF checkpoint; "
f"select the pipeline repo."
)
# Non-GGUF loads fetch + deserialise weights, so gate to unsloth/ or a local path.
if kind != "gguf" and not _is_trusted_diffusion_repo(repo_id):
raise ValueError(
f"Non-GGUF diffusion loads are restricted to unsloth/* repos (or a local "
f"path); got '{repo_id}'. Pass a gguf_filename to load a GGUF instead."
)
# The companion base repo also loads via from_pretrained, so it must clear the same trust bar.
if base_repo and base_repo.strip() and not _is_trusted_diffusion_repo(base_repo):
raise ValueError(
f"base_repo is restricted to unsloth/* repos (or a local path); got "
f"'{base_repo}'."
)
# A local base_repo loads as a full pipeline; reject a non-pipeline one before eviction.
_assert_local_base_is_pipeline(base_repo)
# Reject a bad LOCAL pick before the route evicts chat: a path-shaped repo_id must be on disk.
local_root = Path(repo_id).expanduser()
# Path-shaped: "."/".." prefix, a backslash (never in "org/name"), or an absolute path.
path_shaped = (
repo_id.startswith(("/", "\\", "~", ".")) or "\\" in repo_id or local_root.is_absolute()
)
if kind in ("gguf", "single_file"):
if not gguf_filename:
raise ValueError(f"a single-file checkpoint name is required for a '{kind}' load.")
# Fail a kind/extension mismatch before the handoff: gguf needs .gguf, single_file must not.
is_gguf_name = gguf_filename.lower().endswith(".gguf")
if kind == "gguf" and not is_gguf_name:
raise ValueError("a 'gguf' load requires a .gguf checkpoint name.")
if kind == "single_file" and is_gguf_name:
raise ValueError("a .gguf checkpoint needs model_kind 'gguf', not 'single_file'.")
# A single-file load must name a real .safetensors, else it evicts chat then fails in background.
if kind == "single_file" and not gguf_filename.lower().endswith(".safetensors"):
raise ValueError(
f"'{gguf_filename}' is not a loadable single-file checkpoint "
f"(expected a .safetensors name; use a .gguf name for a GGUF load)."
)
if local_root.exists():
resolve_local_gguf_child(local_root, gguf_filename)
elif path_shaped:
raise FileNotFoundError(f"Local model path does not exist: {repo_id}")
else: # pipeline
if gguf_filename:
raise ValueError(
"a 'pipeline' load takes a full diffusers repo, not a single-file name."
)
if local_root.exists():
if not (local_root / "model_index.json").exists():
raise FileNotFoundError(
f"Local pipeline directory has no model_index.json: {repo_id}"
)
elif path_shaped:
raise FileNotFoundError(f"Local model path does not exist: {repo_id}")
elif repo_id.upper().endswith("-GGUF"):
# A remote "*-GGUF" id is not a pipeline; reject here instead of evicting chat then failing.
raise ValueError(
f"'{repo_id}' is a single-file GGUF repo; load it with model_kind 'gguf' "
f"and a .gguf filename, not as a full pipeline."
)
return fam
# ── Background load + progress ─────────────────────────────────────────
def begin_load(
self,
repo_id: str,
*,
gguf_filename: Optional[str] = None,
base_repo: Optional[str] = None,
family_override: Optional[str] = None,
hf_token: Optional[str] = None,
cpu_offload: bool = False,
memory_mode: Optional[str] = None,
speed_mode: Optional[str] = None,
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
transformer_quant_fast_accum: Optional[bool] = None,
transformer_prequant_path: Optional[str] = None,
attention_backend: Optional[str] = None,
transformer_cache: Optional[str] = None,
transformer_cache_threshold: Optional[float] = None,
model_kind: Optional[str] = None,
loras: Optional[list[tuple[str, float]]] = None,
) -> dict[str, Any]:
"""Validate, then run the (slow) load on a daemon thread. Returns at once."""
# A blank token must mean "anonymous", not an empty credential the Hub 401s.
hf_token = (hf_token.strip() if isinstance(hf_token, str) else hf_token) or None
# base_repo is gated at the route pre-eviction; this only cheap-fails the resolved repo/family.
fam = self.validate_load_request(
repo_id,
gguf_filename = gguf_filename,
family_override = family_override,
model_kind = model_kind,
)
with self._lock:
# Allow starting over a previously-failed load, but not over a live one.
if self._loading is not None and self._loading.error is None:
raise RuntimeError("A diffusion load is already in progress.")
self._load_token += 1
token = self._load_token
# Best-effort download preemption; the token is the real commit guard.
self._cancel_event.clear()
# Seed with the family fallback; the worker resolves the real base and updates this.
self._loading = _LoadingState(repo_id = repo_id, base_repo = fam.base_repo)
threading.Thread(
target = self._run_load,
kwargs = dict(
repo_id = repo_id,
gguf_filename = gguf_filename,
base_repo = base_repo,
family_override = family_override,
hf_token = hf_token,
cpu_offload = cpu_offload,
memory_mode = memory_mode,
speed_mode = speed_mode,
text_encoder_quant = text_encoder_quant,
transformer_quant = transformer_quant,
transformer_quant_fast_accum = transformer_quant_fast_accum,
transformer_prequant_path = transformer_prequant_path,
attention_backend = attention_backend,
transformer_cache = transformer_cache,
transformer_cache_threshold = transformer_cache_threshold,
model_kind = model_kind,
loras = loras,
_load_token = token,
),
daemon = True,
).start()
return self.status()
def _run_load(self, **kwargs: Any) -> None:
token = kwargs.get("_load_token")
try:
# Resolve the base repo and estimate sizes here (both network) so begin_load returns instantly.
fam = detect_family_for_pick(
kwargs["repo_id"], kwargs.get("gguf_filename"), kwargs.get("family_override")
)
kind = resolve_model_kind(kwargs.get("gguf_filename"), kwargs.get("model_kind"))
if kind == "pipeline":
# The full pipeline IS the repo, so the base repo is the repo itself.
base = kwargs["repo_id"]
else:
base = _resolve_base_repo(
kwargs["repo_id"], kwargs.get("base_repo"), fam, kwargs.get("hf_token")
)
kwargs["base_repo"] = base
# The pre-cast encoder injection replaces these weights, so do not prefetch their dense shards. Same resolver as the injection, so prefetch and load never disagree.
te_prequant_files = self._te_prequant_plan_files(
fam, kwargs.get("text_encoder_quant"), kwargs.get("hf_token")
)
expected, base_files = self._estimate_download_bytes(
kwargs["repo_id"],
kwargs.get("gguf_filename"),
base,
kwargs.get("hf_token"),
kind = kind,
single_file_is_pipeline = bool(fam and fam.single_file_is_pipeline),
# Pull the base shards here rather than inside the locked, unpreemptable finalize.
include_transformer = kind == "gguf"
and self._dense_quant_prefetch_needed(fam, kwargs),
skip_te_components = tuple(te_prequant_files),
)
with self._lock:
# Stamp progress only if this load is still current (a superseder has its own token).
if self._load_token == token and self._loading is not None:
self._loading.base_repo = base
self._loading.expected_bytes = expected
# Download outside the lock so unload/an eviction can preempt the pull.
kwargs["_base_local_dir"] = self._prefetch_files(
kwargs["repo_id"],
kwargs.get("gguf_filename"),
base,
base_files,
kwargs.get("hf_token"),
)
self.load_pipeline(**kwargs)
with self._lock:
# Only clear the marker if this load is still current (a superseder has its own token).
if self._load_token == token:
self._loading = None
except Exception as exc: # noqa: BLE001 — surfaced to the client via load_progress
# A cancelled/superseded load raised below; don't log/stamp it onto the current load.
if self._load_token != token:
return
logger.error("diffusion.load_failed: %s", exc)
# Free the debris of a failed construction; guarded so a sticky CUDA error still stamps the error.
try:
clear_gpu_cache()
except Exception: # noqa: BLE001
pass
# Redact native paths: this error is surfaced verbatim and Studio can be shared.
from utils.native_path_leases import redact_native_paths
with self._lock:
if self._load_token == token and self._loading is not None:
self._loading.error = redact_native_paths(str(exc))
def load_progress(self) -> dict[str, Any]:
"""Phase + downloaded/total bytes for the in-flight load (cache-scan based)."""
loading = self._loading
if loading is not None and loading.error:
return _progress("error", error = loading.error)
if loading is None:
return _progress("ready" if self._state is not None else None)
# Sum checkpoint + companion cache; for a full-pipeline load base IS the repo, so count once.
downloaded = self._cache_bytes(loading.repo_id)
if loading.base_repo and loading.base_repo != loading.repo_id:
downloaded += self._cache_bytes(loading.base_repo)
expected = loading.expected_bytes
# Downloads done, still finalizing. The cache scan can exceed the estimate, so clamp to 100%.
if expected > 0 and downloaded >= expected * 0.999:
return _progress("finalizing", min(downloaded, expected), expected, 1.0)
if expected <= 0:
# No size estimate (lookup failed, or everything was already cached): report the phase with no byte claim, since `downloaded` scans what is PRESENT, not what this load fetched.
return _progress("downloading")
return _progress("downloading", downloaded, expected, min(downloaded / expected, 1.0))
def loading_repo_ids(self) -> tuple[str, ...]:
"""Repo ids an in-flight background load is downloading (empty when idle).
The delete-cached guard needs this: during a load ``status()["loaded"]`` is
still False, but deleting the target repo (or its companion base) would yank
blobs and snapshot files from under the download/assembly."""
with self._lock:
loading = self._loading
if loading is None or loading.error is not None:
return ()
return tuple(r for r in (loading.repo_id, loading.base_repo) if r)
@staticmethod
def _te_prequant_plan_files(
fam: Any, text_encoder_quant: Optional[str], hf_token: Optional[str]
) -> dict[str, tuple[str, list[tuple[str, int]]]]:
"""``{component: (repo_id, [(rfilename, size)])}`` for the text encoders this pick will
take PRE-CAST from a hosted checkpoint instead of the base repo's dense weights.
Empty unless the request asked for a scheme with a hosted artifact AND that artifact
really resolves, so a plan can never drop a dense encoder the load still wants."""
try:
from huggingface_hub import HfApi
from .diffusion_te_prequant import te_prequant_hub_files, te_prequant_sources
sources = te_prequant_sources(
fam,
te_quant_mode = text_encoder_quant,
target = resolve_diffusion_device_target(),
)
if not sources:
return {}
files = te_prequant_hub_files(sources, HfApi(token = hf_token or None), logger)
return {c: (sources[c].location, f) for c, f in files.items()}
except Exception as exc: # noqa: BLE001 -- an unresolvable pre-cast keeps the dense encoder
logger.warning("diffusion.te_prequant_plan_failed: %s", exc)
return {}
@staticmethod
def _estimate_download_bytes(
repo_id: str,
gguf_filename: Optional[str],
base_repo: str,
hf_token: Optional[str],
*,
kind: str = "gguf",
single_file_is_pipeline: bool = False,
include_transformer: bool = False,
sizes_out: Optional[dict[str, int]] = None,
skip_te_components: tuple[str, ...] = (),
) -> tuple[int, list[str]]:
"""Total download size for the progress bar, plus the base-repo files to
fetch (the prefetch reuses this list, so the base is listed only once).
``sizes_out``, when given, is filled with per-repo byte totals so the download
plan can size one job per repo off this same single pair of Hub lookups.
For a ``pipeline`` load the whole repo IS the pipeline (``base_repo`` is the
repo itself), so the transformer/ subfolder is INCLUDED -- unlike the GGUF /
single-file paths, where the transformer is the single file and the base repo
supplies only the companions. For a ``single_file_is_pipeline`` family (SDXL) the
single file is the WHOLE pipeline, so the base repo supplies only config/tokenizer
(no weights) and its weight files are skipped.
``skip_te_components`` names the text encoders this pick loads PRE-CAST from a hosted
checkpoint, so their dense weight shards are not counted or fetched: staging the dense
encoder for a pre-cast load wastes tens of GB (FLUX.2-dev's Mistral-24B is ~48 GB,
Qwen-Image's Qwen2.5-VL ~16.6 GB) and nothing ever opens them. Everything else in the
component folder (config, shard index, tokenizer) is kept -- the pre-cast loader
meta-inits the encoder from the base repo's config."""
from huggingface_hub import HfApi
from .diffusion_te_prequant import is_prequant_covered_weight
api = HfApi()
total = 0
base_files: list[str] = []
def _dense_te_shard(rfilename: str) -> bool:
return bool(skip_te_components) and is_prequant_covered_weight(
rfilename, skip_te_components
)
try:
if kind == "pipeline":
info = api.model_info(repo_id, files_metadata = True, token = hf_token)
picked = [
s
for s in info.siblings
if _pipeline_file_downloaded(s.rfilename) and not _dense_te_shard(s.rfilename)
]
# diffusers prefers safetensors: drop a .bin whose dir also has a picked .safetensors.
st_dirs = {
s.rfilename.rsplit("/", 1)[0]
for s in picked
if s.rfilename.endswith(".safetensors")
}
for s in picked:
if s.rfilename.endswith(".bin") and s.rfilename.rsplit("/", 1)[0] in st_dirs:
continue
base_files.append(s.rfilename)
total += s.size or 0
if sizes_out is not None:
sizes_out[repo_id] = total
return total, base_files
# Skip the Hub size lookup for a LOCAL gguf path: model_info raises on a filesystem path, and the catch below would force a synchronous companion pull.
if gguf_filename and not Path(repo_id).expanduser().exists():
info = api.model_info(repo_id, files_metadata = True, token = hf_token)
gguf_bytes = sum(s.size or 0 for s in info.siblings if s.rfilename == gguf_filename)
total += gguf_bytes
if sizes_out is not None:
sizes_out[repo_id] = gguf_bytes
# A whole-pipeline single file (SDXL) needs only the base's config/tokenizer, not its weights.
if kind == "single_file" and single_file_is_pipeline:
base_filter = _base_config_file_downloaded
else:
def base_filter(rfilename: str) -> bool:
return _base_file_downloaded(rfilename, include_transformer = include_transformer)
base_info = api.model_info(base_repo, files_metadata = True, token = hf_token)
base_bytes = 0
for s in base_info.siblings:
if base_filter(s.rfilename) and not _dense_te_shard(s.rfilename):
base_files.append(s.rfilename)
base_bytes += s.size or 0
total += base_bytes
if sizes_out is not None:
sizes_out[base_repo] = base_bytes
except Exception as exc: # noqa: BLE001 — estimate is best-effort
logger.warning("diffusion.size_estimate_failed: %s", exc)
return total, base_files
def download_plan(
self,
repo_id: str,
*,
gguf_filename: Optional[str] = None,
base_repo: Optional[str] = None,
family_override: Optional[str] = None,
model_kind: Optional[str] = None,
hf_token: Optional[str] = None,
text_encoder_quant: Optional[str] = None,
**load_kwargs: Any,
) -> dict[str, Any]:
"""The repos + exact files this pick needs, so the Hub download manager can fetch
them with the same file scope the loader would.
A plain snapshot_download would also pull what the loader deliberately skips (the
packaged root single, transformer/ shards, fp16 twins) -- tens of GB per FLUX repo.
Resolves family/kind/base exactly as ``_run_load`` does, so the plan and the load
agree. Local paths are already on disk and yield no entries.
``text_encoder_quant`` is read for the same reason as the DiT quant: an fp8 request
loads a hosted PRE-CAST encoder, so the base repo's dense encoder shards must not be
staged and the pre-cast checkpoint must be. Without it the manager stages the dense
encoder (tens of GB the load never opens) and the load then pulls the pre-cast file
inline, outside the manager's progress and disk preflight."""
fam = detect_family_for_pick(repo_id, gguf_filename, family_override)
kind = resolve_model_kind(gguf_filename, model_kind)
if kind == "pipeline":
base = repo_id # the full pipeline IS the repo
else:
base = _resolve_base_repo(repo_id, base_repo, fam, hf_token)
# Only a checkpoint that really resolves on the Hub earns the right to drop dense shards.
te_files = self._te_prequant_plan_files(fam, text_encoder_quant, hf_token)
sizes: dict[str, int] = {}
total, base_files = self._estimate_download_bytes(
repo_id,
gguf_filename,
base,
hf_token,
kind = kind,
single_file_is_pipeline = bool(fam and fam.single_file_is_pipeline),
include_transformer = kind == "gguf"
and self._dense_quant_prefetch_needed(fam, load_kwargs),
sizes_out = sizes,
skip_te_components = tuple(te_files),
)
entries: list[dict[str, Any]] = []
for repo, files in te_files.values():
entries.append(
{
"repo_id": repo,
"files": [name for name, _size in files],
"bytes": int(sum(size for _name, size in files)),
"gguf_filename": None,
}
)
total += int(sum(size for _name, size in files))
if gguf_filename and not Path(repo_id).expanduser().exists():
entries.append(
{
"repo_id": repo_id,
"files": [gguf_filename],
"bytes": int(sizes.get(repo_id, 0)),
"gguf_filename": gguf_filename,
}
)
if base_files and not Path(base).expanduser().exists():
entries.append(
{
"repo_id": base,
"files": base_files,
"bytes": int(sizes.get(base, 0)),
"gguf_filename": None,
}
)
return {"entries": entries, "total_bytes": int(total)}
@staticmethod
def _hub_cache_repo_dir(repo_id: str) -> Path:
"""Local HF hub cache dir for ``repo_id``.
Reads the live setting, not huggingface_hub's import-time constant: changing the
cache folder does not update the constant, so the old one would count bytes in a
root the download no longer writes to (progress stuck at 0 for the whole pull)."""
return Path(hub_cache_dir()) / f"models--{repo_id.replace('/', '--')}"
@staticmethod
def _cache_bytes(repo_id: str) -> int:
blobs = DiffusionBackend._hub_cache_repo_dir(repo_id) / "blobs"
total = 0
try:
for entry in blobs.iterdir():
try:
total += entry.stat().st_size
except OSError:
continue # broken symlink / unreadable
except OSError:
return 0 # repo not in cache yet
return total
@staticmethod
def _local_dir_weight_bytes(path: Path, *, exclude_transformer: bool) -> int:
"""Sum the on-disk weight files under a local diffusers directory. The HF blob
cache is empty for a local path, so this is the only size signal for auto memory
planning; without it a large local model folds to zero and the planner skips
offload and OOMs. ``exclude_transformer`` drops the ``transformer/`` subfolder
for GGUF/single-file loads (their transformer is the single file, not resident
here); a full pipeline load keeps it (the whole repo is resident)."""
total = 0
for f in path.rglob("*"):
if f.suffix.lower() not in (".safetensors", ".bin", ".pt", ".ckpt"):
continue
try:
rel = f.relative_to(path)
except ValueError:
continue
if exclude_transformer and rel.parts and rel.parts[0] == "transformer":
continue
try:
total += f.stat().st_size
except OSError:
continue
return total
@staticmethod
def _max_over_cached_revs(base: str, fn: Callable[[Path], int]) -> int:
"""Apply ``fn`` to a LOCAL diffusers dir, or to the fullest cached hub snapshot
revision (the active one is the fullest), returning that count. 0 when nothing is
cached. Multiple revisions may be cached, so take the max."""
local = Path(base).expanduser()
if local.is_dir():
return fn(local)
snapshots = DiffusionBackend._hub_cache_repo_dir(base) / "snapshots"
if not snapshots.is_dir():
return 0
return max((fn(rev) for rev in snapshots.iterdir() if rev.is_dir()), default = 0)
@staticmethod
def _companion_cache_bytes(base: str) -> int:
"""Resident companion (VAE + text-encoder) size for the memory plan.
Excludes ``transformer/`` (supplied by the GGUF/single file, not resident here) --
otherwise the dense-quant prefetch's cached transformer shards would inflate this
and wrongly force offload. Walks the snapshot dir, not the flat ``blobs/`` cache,
since only the snapshot preserves the subfolder split needed to exclude it."""
return DiffusionBackend._max_over_cached_revs(
base, lambda d: DiffusionBackend._local_dir_weight_bytes(d, exclude_transformer = True)
)
@staticmethod
def _safetensors_param_count(path: Path) -> int:
"""Total tensor elements in a safetensors file, read from its JSON header without
touching the tensor data. 0 on any read/parse failure."""
try:
with open(path, "rb") as fh:
header_len = int.from_bytes(fh.read(8), "little")
header = json.loads(fh.read(header_len))
total = 0
for name, meta in header.items():
if name == "__metadata__" or not isinstance(meta, dict):
continue
numel = 1
for dim in meta.get("shape", []):
numel *= dim
total += numel
return total
except Exception: # noqa: BLE001 — corrupt/crafted shard degrades to 0, never crashes the load
return 0
@staticmethod
def _dense_transformer_resident_bytes(base: str) -> int:
"""Resident bf16 size of the base repo's dense ``transformer/`` for the dense-quant
preflight. That fast path loads the transformer at the compute dtype (bf16, 2
bytes/param) before quantizing, so budget num_params * 2 -- NOT the on-disk bytes,
which for an F32 base (e.g. Z-Image) are ~2x the resident size. Read from the
safetensors shard headers. Returns 0 when no ``transformer/*.safetensors`` shards
are present (an uncached base, or a .bin-only transformer); the caller then gates
the fast path on the plain plan."""
def _params(d: Path) -> int:
tdir = d / "transformer"
if not tdir.is_dir():
return 0
return sum(
DiffusionBackend._safetensors_param_count(s) for s in tdir.glob("*.safetensors")
)
return DiffusionBackend._max_over_cached_revs(base, _params) * 2 # bf16: 2 bytes/param
# ── Synchronous load / generate / unload ───────────────────────────────
def load_pipeline(
self,
repo_id: str,
*,
gguf_filename: Optional[str] = None,
base_repo: Optional[str] = None,
family_override: Optional[str] = None,
hf_token: Optional[str] = None,
cpu_offload: bool = False,
memory_mode: Optional[str] = None,
speed_mode: Optional[str] = None,
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
transformer_quant_fast_accum: Optional[bool] = None,
transformer_prequant_path: Optional[str] = None,
attention_backend: Optional[str] = None,
transformer_cache: Optional[str] = None,
transformer_cache_threshold: Optional[float] = None,
model_kind: Optional[str] = None,
# LoRA adapters to BAKE into a torchao int8/fp8 build. Ignored elsewhere: bf16/bnb take adapters at generation time, GGUF-as-is has no dense transformer.
loras: Optional[list[tuple[str, float]]] = None,
_load_token: Optional[int] = None,
_base_local_dir: Optional[str] = None,
) -> dict[str, Any]:
# A blank token must degrade to anonymous, not be passed as a credential. Normalize once.
hf_token = hf_token.strip() if isinstance(hf_token, str) else hf_token
hf_token = hf_token or None
# Validate first (no torch/diffusers) so a bad family fails even in a no-diffusers runtime.
hf_token = (hf_token.strip() if isinstance(hf_token, str) else hf_token) or None
# base_repo is gated at the route pre-eviction; this only cheap-fails the resolved repo/family.
fam = self.validate_load_request(
repo_id,
gguf_filename = gguf_filename,
family_override = family_override,
model_kind = model_kind,
)
kind = resolve_model_kind(gguf_filename, model_kind)
# Validate every mode string that can raise BEFORE this load evicts the previous pipeline (transformer_quant validate-only, keeping the unset/auto vs explicit-off tri-state).
normalize_transformer_quant(transformer_quant)
normalize_speed_mode(speed_mode)
normalize_attention_backend(attention_backend)
normalize_transformer_cache(transformer_cache)
normalize_te_quant(text_encoder_quant)
# A full pipeline is its own base; single-file kinds resolve the companion base repo.
base = (
repo_id if kind == "pipeline" else _resolve_base_repo(repo_id, base_repo, fam, hf_token)
)
target = self._resolve_device_target(fam)
device, dtype = target.device, target.dtype
import diffusers
# Pre-install the optional attention kernel before the load locks: the pip install can take 600s and would block unload/cancel. Best-effort; the locked resolve is a no-op.
try:
preinstall_backend = select_attention_backend(
target, attention_backend, speed_active = True
)
if preinstall_backend is not None:
_ensure_attention_backend_installed(preinstall_backend, logger)
except Exception: # noqa: BLE001 — the locked path re-resolves and validates
pass
# Abort an in-flight denoise and wait for it to exit: a load claims VRAM, so it must not overlap.
with self._lock:
# Bail before signalling if this load was superseded, else a stale worker aborts a live one.
if _load_token is not None and _load_token != self._load_token:
raise RuntimeError("Diffusion load was cancelled.")
if self._active_generate_cancel is not None:
self._active_generate_cancel.set()
# Same fence unload() takes: a generation queued behind the active denoise would
# otherwise slip in here and run on the pipeline this load is about to free.
self._teardown_waiters += 1
with self._generate_lock:
with self._lock:
try:
# Re-check: a newer load/unload may have superseded this one while we waited.
if _load_token is not None and _load_token != self._load_token:
raise RuntimeError("Diffusion load was cancelled.")
# Free the old pipeline before allocating the new one (never two in VRAM).
self._unload_locked()
finally:
# Released here, not at the end of the load: the old pipe is gone (or this load
# bailed), and the rest of the load holds _generate_lock anyway.
self._teardown_waiters -= 1
# Single-file kinds resolve a checkpoint path; the pipeline kind has none.
single_file_path = (
self._resolve_gguf_path(repo_id, gguf_filename, hf_token)
if kind in ("gguf", "single_file")
else None
)
transformer_cls = getattr(diffusers, fam.transformer_class)
pipeline_cls = getattr(diffusers, fam.pipeline_class)
# Decide placement up front (weights still on CPU). Budgets the GGUF; dense is preflighted below.
plan = self._plan_memory(
target,
single_file_path,
base,
fam,
memory_mode,
cpu_offload,
kind = kind,
repo_id = repo_id,
)
# Dtype tri-state: unset/"auto" -> hardware ladder picks a quantised build; "none"/"off" pins GGUF-as-is; an explicit scheme pins it. An overwritten "auto" still records source=auto.
if transformer_quant is None or str(transformer_quant).strip().lower() in (
"",
"auto",
):
# An explicit Speed="off" stays GGUF-as-is: auto-quant would break the bit-exact request.
speed_off = (
speed_mode is not None and str(speed_mode).strip().lower() == SPEED_OFF
)
transformer_quant = "off" if speed_off else TQ_AUTO
# Default-on fast path (GGUF kind only): load the dense bf16 transformer and torchao-quantise it, which beats
# GGUF per-matmul dequant on speed AND quality. Needs CUDA + bf16 + a resident fit; any failure falls back to the GGUF build.
pipe = None
transformer_quant_engaged = None
quant_plan = None
# The GGUF-size plan can mis-budget the fast path, so preflight the real footprint pre-eviction.
dense_declined = False
# False when the plan only holds a PREQUANT-sized build: a failed prequant load must raise, not materialise the unbudgeted dense bf16 transformer.
dense_fallback_allowed = True
if (
kind == "gguf"
and normalize_transformer_quant(transformer_quant) is not None
and dense_transformer_supported(target)
):
if plan.offload_policy != OFFLOAD_NONE:
# The GGUF plan picked offload but the quantised artifact is smaller, so re-plan against the candidate: a resident quant build beats an offloaded GGUF.
candidate = resolve_dense_quant_candidate(
fam = fam,
target = target,
requested = transformer_quant,
base_repo = base,
prequant_path = transformer_prequant_path,
# A LoRA bake skips the prequant shortcut, so size for the dense build it will run.
force_dense = bool(loras),
logger = logger,
)
if candidate is not None:
def _replan_candidate():
return self._plan_memory(
target,
single_file_path,
base,
fam,
memory_mode,
cpu_offload,
kind = kind,
repo_id = repo_id,
transformer_resident_override_mib = (
candidate.transient_transformer_mib
),
# Pass the companion estimate so prefetched base shards aren't double-counted.
companion_override_mib = candidate.companions_mib,
)
replanned = _replan_candidate()
if (
replanned.offload_policy != OFFLOAD_NONE
# Explicit balanced/low_vram picks offload BY MODE, so a fresh snapshot cannot change it.
and normalize_memory_mode(memory_mode)
not in (MEMORY_MODE_BALANCED, MEMORY_MODE_LOW_VRAM)
and plan_fits_total_capacity(replanned)
):
# The candidate fits TOTAL capacity but the instantaneous free reading said no; a transient foreign allocation (~100 GB briefly on an idle B200) must not force the GGUF fallback, so re-snapshot (settled) and replan once.
replanned = _replan_candidate()
if replanned.offload_policy != OFFLOAD_NONE:
logger.info(
"diffusion.transformer_quant_declined: required=%s MiB "
"budget=%s MiB free=%s MiB policy=%s (%s)",
replanned.estimates.get("resident_required_mib"),
replanned.estimates.get("safe_device_budget_mib"),
getattr(replanned.device_memory, "free_mib", None),
replanned.offload_policy,
"; ".join(replanned.reasons),
)
if replanned.offload_policy == OFFLOAD_NONE:
quant_plan = replanned
# The GGUF plan declined resident; a prequant-sized replan says nothing about the dense build.
if candidate.prequant:
dense_fallback_allowed = False
else:
# This path materialises the dense bf16 transformer (bigger than the GGUF), so re-check the fit rather than OOMing after eviction. Skipped for a prequant, which loads a small file.
scheme = select_transformer_quant_scheme(
target,
transformer_quant, # normalized above
family = getattr(fam, "name", None),
)
# usable_prequant_source (not resolve_): a missing/non-allowlisted local path must not count here, or the dense-fit re-check is skipped and materialising OOMs.
prequant = (
# A LoRA bake skips the prequant shortcut, so gate the fast path as if no prequant existed.
None
if loras
else usable_prequant_source(
fam,
scheme,
path_override = transformer_prequant_path,
base_repo = base,
)
if scheme is not None
else None
)
dense_mib = int(
self._dense_transformer_resident_bytes(base) // (1024 * 1024)
)
if dense_mib > 0:
dense_plan = self._plan_memory(
target,
single_file_path,
base,
fam,
memory_mode,
cpu_offload,
kind = kind,
repo_id = repo_id,
transformer_resident_override_mib = dense_mib,
)
if dense_plan.offload_policy != OFFLOAD_NONE:
dense_fallback_allowed = False
# Without a prequant source the dense build is the only path, so a dense misfit skips the fast path; with one, only the dense fallback is forbidden.
if prequant is None:
dense_declined = True
if (
kind == "gguf"
and normalize_transformer_quant(transformer_quant) is not None
and dense_transformer_supported(target)
and not dense_declined
and (plan.offload_policy == OFFLOAD_NONE or quant_plan is not None)
):
try:
pipe, transformer_quant_engaged = self._load_dense_quant_pipeline(
transformer_cls,
pipeline_cls,
base,
device,
dtype,
hf_token,
target,
transformer_quant,
transformer_quant_fast_accum,
fam = fam,
base_local_dir = _base_local_dir,
prequant_path = transformer_prequant_path,
allow_dense_fallback = dense_fallback_allowed,
lora_specs = loras,
text_encoder_quant = text_encoder_quant,
)
except Exception as exc: # noqa: BLE001 — fall back to the GGUF build
logger.warning(
"diffusion.transformer_quant_fallback: %s (loading GGUF)", exc
)
pipe = None
transformer_quant_engaged = None
# Drop the exception before clearing the cache: its traceback pins the dense transformer's VRAM.
del exc
# Guarded: a sticky CUDA error can raise; the fallback must reach the GGUF build.
try:
clear_gpu_cache()
except Exception: # noqa: BLE001
pass
if transformer_quant_engaged is not None and quant_plan is not None:
# The engaged dense build uses the re-planned placement; the GGUF-size plan stays for fallback.
plan = quant_plan
if (
pipe is None
and kind == "gguf"
and normalize_transformer_quant(transformer_quant) is not None
and any(w != 0 for (_lid, w) in (loras or ()))
):
# Adapters were requested BAKED but that build was declined or failed, and the GGUF fallback cannot carry them; fail loudly rather than silently dropping the LoRAs.
raise RuntimeError(
"The requested LoRA adapters could not be applied: baking adapters "
"requires the quantized (int8/fp8) transformer build, which was "
"declined or failed on this device (see the server log), and the "
"GGUF fallback cannot carry them. Retry without transformer_quant "
"adapters, free VRAM, or pick a smaller model."
)
if pipe is None:
if kind == "pipeline":
# Full diffusers repo: from_pretrained pulls every component and re-applies embedded quant config.
if fam.name == KREA2_FAMILY_NAME:
# krea ships transformers-5.x configs the 4.x line cannot parse, so assemble per-component (diffusion_krea2.py). That path never sees pipe_kwargs, so pass the pre-cast TE.
pipe = load_krea2_pipeline(
repo_id,
dtype,
hf_token = hf_token,
text_encoder = te_prequant_pipe_kwargs(
fam,
repo_id,
te_quant_mode = text_encoder_quant,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
).get("text_encoder"),
)
elif fam.name == IDEOGRAM4_FAMILY_NAME:
# ideogram ships the same transformers-5.x Qwen stack as krea; assemble per-component too.
pipe = load_ideogram4_pipeline(repo_id, dtype, hf_token = hf_token)
else:
pipe_kwargs: dict[str, Any] = {
"torch_dtype": dtype,
"cache_dir": hub_cache_dir(),
}
if hf_token:
pipe_kwargs["token"] = hf_token
if fam.name == HIDREAM_FAMILY_NAME:
# The repo names a Llama text_encoder_4 it does not ship; supply it from the open mirror.
pipe_kwargs.update(
hidream_te4_kwargs(
dtype,
hf_token,
fam = fam,
te_quant_mode = text_encoder_quant,
target = target,
)
)
# A hosted pre-cast fp8 text encoder skips the dense TE download; quantize_text_encoders re-applies the cast idempotently.
pipe_kwargs.update(
te_prequant_pipe_kwargs(
fam,
repo_id,
te_quant_mode = text_encoder_quant,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
)
)
# The prefetched snapshot dir keeps from_pretrained off the hub (24 GB per FLUX.1 otherwise).
pipe = pipeline_cls.from_pretrained(
_base_local_dir or repo_id, **pipe_kwargs
)
elif kind == "single_file" and fam.single_file_is_pipeline:
# A single-file SDXL-style checkpoint is the WHOLE pipeline: load it through the pipeline class, with ``config`` pointing at the base repo for the structure.
sf_pipe_kwargs: dict[str, Any] = {
"torch_dtype": dtype,
"config": base,
"cache_dir": hub_cache_dir(),
}
if hf_token:
sf_pipe_kwargs["token"] = hf_token
pipe = pipeline_cls.from_single_file(single_file_path, **sf_pipe_kwargs)
else:
# Transformer-only single file; VAE/text-encoder/scheduler come from the base repo.
sf_kwargs: dict[str, Any] = {
"torch_dtype": dtype,
"config": base,
"subfolder": "transformer",
# Config is fetched from the (possibly gated) base before auth.
"token": hf_token,
"cache_dir": hub_cache_dir(),
}
if kind == "gguf":
# Dequantise the GGUF transformer on-device at the compute dtype.
sf_kwargs["quantization_config"] = diffusers.GGUFQuantizationConfig(
compute_dtype = dtype
)
# sd.cpp GGUFs prefix tensors with model.diffusion_model.; the FLUX.2 / Qwen converters choke.
_install_gguf_prefix_strip(transformer_cls, logger)
# A safetensors single-file (fp8) carries its own dtype: no GGUF dequant config.
transformer = transformer_cls.from_single_file(
single_file_path, **sf_kwargs
)
if fam.name == KREA2_FAMILY_NAME:
pipe = load_krea2_pipeline(
base,
dtype,
hf_token = hf_token,
transformer = transformer,
# Same pre-cast TE hand-in as the full-pipeline branch.
text_encoder = te_prequant_pipe_kwargs(
fam,
base,
te_quant_mode = text_encoder_quant,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
).get("text_encoder"),
)
else:
pipe_kwargs = {
"torch_dtype": dtype,
"transformer": transformer,
"cache_dir": hub_cache_dir(),
}
if hf_token:
pipe_kwargs["token"] = hf_token
if fam.name == HIDREAM_FAMILY_NAME:
# Same Llama TE4 assembly as the full-pipeline branch above.
pipe_kwargs.update(
hidream_te4_kwargs(
dtype,
hf_token,
fam = fam,
te_quant_mode = text_encoder_quant,
target = target,
)
)
# Same pre-cast TE injection as above: the GGUF supplies the transformer, so the TE is the big download.
pipe_kwargs.update(
te_prequant_pipe_kwargs(
fam,
base,
te_quant_mode = text_encoder_quant,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
)
)
pipe = pipeline_cls.from_pretrained(
_base_local_dir or base, **pipe_kwargs
)
# Effective speed: GGUF defaults to near-lossless `default` (~2.2x, below the quant noise floor); dense stays bit-identical `off`. Explicit is honored.
effective_speed = resolve_speed_mode(speed_mode, is_gguf = kind == "gguf")
# A torchao-quantized dense transformer must be compiled (eager is ~30x slower, losing to GGUF).
if transformer_quant_engaged is not None and effective_speed == SPEED_OFF:
logger.info(
"diffusion.transformer_quant: forcing speed_mode=default "
"(quantized transformer must be compiled; eager is ~30x slower)"
)
effective_speed = SPEED_DEFAULT
# Deferred speed auto for dense models: stay eager, but generate() engages `default` on the 3rd image, where repeated use amortises the compile. Only when speed was unset.
speed_deferred = (
speed_mode is None
and effective_speed == SPEED_OFF
and transformer_quant_engaged is None
and compile_eligible(target, is_gguf = False, family = fam)
)
# Speed optims run BEFORE placement (channels_last/compile precede offload), so snapshot the global backend flags first for unload restore.
backend_flags_before = snapshot_backend_flags()
# Pick the attention kernel BEFORE compile: auto upgrades to cuDNN fused attention on NVIDIA when a speed profile is active (~1.18x); explicit is honored.
attention_engaged = apply_attention_backend(
pipe,
select_attention_backend(
target, attention_backend, speed_active = effective_speed != SPEED_OFF
),
logger = logger,
)
# Step caching (First-Block-Cache), also before compile: reuses the transformer tail across steps (~1.4x on Flux at LPIPS ~0.08) and drops compile fullgraph when engaged.
# Tri-state: unset/"auto" -> step-count policy decides (FBCACHE_MIN_STEPS); "off"/"fbcache" pinned.
cache_request = normalize_transformer_cache(transformer_cache)
cache_auto = transformer_cache is None or cache_request == TC_AUTO
cache_quant_active = transformer_quant_engaged is not None or bool(gguf_filename)
default_steps: Optional[int] = None
if cache_auto:
default_steps, _ = default_generation_params(
gguf_filename, repo_id, base, fam.name
)
cache_request = TC_FBCACHE if default_steps >= FBCACHE_MIN_STEPS else None
cache_engaged = apply_step_cache(
pipe,
mode = cache_request,
threshold = transformer_cache_threshold,
# GGUF transformers are quantized too, so the cache needs the higher threshold.
quant_active = cache_quant_active,
logger = logger,
)
# An auto decision can flip at generation time, but only on a cache-capable transformer.
cache_may_toggle = cache_auto and callable(
getattr(getattr(pipe, "transformer", None), "enable_cache", None)
)
if cache_auto:
if cache_engaged:
cache_reason = (
f"auto: {default_steps}-step default schedule reaches "
f"{FBCACHE_MIN_STEPS}; re-checked per generation"
)
elif cache_request is not None:
cache_reason = "auto: model does not support step caching"
else:
cache_reason = (
f"auto: {default_steps}-step default schedule is below "
f"{FBCACHE_MIN_STEPS}; re-checked per generation"
)
else:
cache_reason = "requested"
# Everything from here to the _LoadState commit mutates PROCESS-WIDE state (class patches, TORCHINDUCTOR_CACHE_DIR, backend flags); the try/finally below restores it on failure.
# gguf_transformer: the dense fast path still sets gguf_filename (fallback) but pipe.transformer is dense (needs REGIONAL block compile), so treat it as non-GGUF.
gguf_transformer = kind == "gguf" and transformer_quant_engaged is None
eager_patched = False
compile_ctx = None
state_committed = False
# Lazy import (these modules import torch) keeps diffusion.py torch-free to import.
from .diffusion_eager_patches import (
install_compile_safe_patches,
uninstall_patches,
)
from .diffusion_arch_patches import (
install_arch_patches,
uninstall_arch_patches,
)
try:
if effective_speed != SPEED_OFF:
install_compile_safe_patches()
# Per-arch compile-safe fusions; neutral under compile, tracked by the same eager_patched flag.
install_arch_patches()
eager_patched = True
else:
uninstall_patches()
uninstall_arch_patches()
# Pre-warmed torch.compile cache: point inductor at a per-fingerprint dir and load a matching bundle before the first compiled forward, so the 25-58s compile is paid once. A miss is silent.
if effective_speed in (SPEED_DEFAULT, SPEED_MAX) and compile_eligible(
target, is_gguf = gguf_transformer, family = fam
):
compile_ctx = compile_cache.begin(
family = fam.name,
# U-Net families (SDXL) carry the denoiser as pipe.unet.
transformer = getattr(pipe, "transformer", None)
or getattr(pipe, "unet", None),
dtype = getattr(target, "dtype", None),
# A GGUF transformer compiles a DIFFERENT graph than a dense load of the same family, so it must key its own bundles.
quant = "gguf" if gguf_transformer else transformer_quant_engaged,
attention_backend = attention_engaged,
compile_kwargs = {
# Mirrors apply_speed_optims fullgraph decision (an active or still-toggleable step cache, or a planned offload, graph-breaks), so the bundle keys on the same setting.
"fullgraph": cache_engaged is None
and not cache_may_toggle
and plan.offload_policy == OFFLOAD_NONE,
"dynamic": effective_speed != SPEED_MAX,
"mode": "max-autotune-no-cudagraphs"
if effective_speed == SPEED_MAX
else "default",
},
logger = logger,
)
speed_applied = apply_speed_optims(
pipe,
target,
is_gguf = gguf_transformer,
family = fam,
speed_mode = effective_speed,
cache_active = cache_engaged is not None or cache_may_toggle,
# Offload installs compiler-disabled onload hooks, so compile drops fullgraph.
offload_active = plan.offload_policy != OFFLOAD_NONE,
logger = logger,
)
if transformer_quant_engaged is not None and not speed_applied.get("compiled"):
# Compile could not engage: the quantized transformer runs eager, far slower than the GGUF it replaced. Surface it loudly.
logger.warning(
"diffusion.transformer_quant: %s engaged but the transformer is NOT "
"compiled; eager torchao quant is ~30x slower than GGUF here",
transformer_quant_engaged,
)
# Quantise the dense companion text encoder(s) before placement so offload moves the smaller weights. Family drives int8 keep-bf16 schedule.
te_quant = quantize_text_encoders(
pipe,
target,
mode = text_encoder_quant,
family = fam.name,
offload_active = plan.offload_policy != OFFLOAD_NONE,
logger = logger,
)
# Persistent conditioning cache (opt-in via UNSLOTH_DIFFUSION_COND_CACHE_DIR): repeated prompts skip the text-encoder forward, so warm prompts never onload the multi-GB encoders under offload.
# Installed AFTER the TE quant so the key reflects the encoders that run; ``base`` keys the companion repo, so the same checkpoint against a different base cannot reuse its embeddings. Dies with the pipe.
cond_cache.install(
pipe,
family = fam.name,
repo_id = repo_id,
base_repo = base,
dtype = dtype,
te_quant = te_quant,
logger = logger,
)
# Apply the planned placement; apply_memory_plan returns the policy/tiling ACTUALLY engaged so status stays honest. Idempotent for the `none` policy.
effective_policy, effective_tiling = apply_memory_plan(
pipe, plan, device = device, logger = logger
)
# Per-control provenance for status. cpu_offload=False is the unset default, so only True is an explicit request.
resolved = build_resolved_record(
{
"speed_mode": (
speed_mode,
"deferred" if speed_deferred else effective_speed,
"quantized transformer requires compile"
if transformer_quant_engaged is not None
and normalize_speed_mode(speed_mode) in (None, SPEED_OFF)
else "auto: exact eager for the first two images; "
"the compile profile engages on the 3rd"
if speed_deferred
else "per-kind default"
if speed_mode is None
else "requested",
),
"transformer_quant": (
transformer_quant,
transformer_quant_engaged or "off",
# The None reason matches the load kind (GGUF loaded vs dense kept).
(
"not engaged (GGUF transformer loaded)"
if kind == "gguf"
else "dense transformer kept unquantized"
)
if transformer_quant_engaged is None
else "re-planned resident for the quantised artifact"
if quant_plan is not None
else "engaged on the dense fast path",
),
"attention_backend": (
attention_backend,
attention_engaged or "native",
"cuDNN fused attention upgrade"
if attention_engaged and attention_backend is None
else "diffusers default"
if attention_engaged is None
else "requested",
),
"memory_mode": (
memory_mode,
effective_policy,
"everything fits on the GPU, no offload needed"
if effective_policy == OFFLOAD_NONE
else "planned from measured free VRAM vs estimated footprint",
),
"transformer_cache": (
None if cache_auto else transformer_cache,
cache_engaged or "off",
cache_reason,
),
"cpu_offload": (
True if cpu_offload else None,
effective_policy != OFFLOAD_NONE,
"legacy flag" if cpu_offload else "from the memory plan",
),
}
)
self._state = _LoadState(
pipe = pipe,
family = fam,
repo_id = repo_id,
base_repo = base,
device = device,
dtype = str(dtype).replace("torch.", ""),
kind = kind,
cpu_offload = effective_policy != OFFLOAD_NONE,
offload_policy = effective_policy,
vae_tiling = effective_tiling,
memory_mode = plan.requested_mode,
speed_mode = effective_speed,
speed_optims = tuple(k for k, v in speed_applied.items() if v),
backend_flags_before = backend_flags_before,
text_encoder_quant = te_quant,
transformer_quant = transformer_quant_engaged,
attention_backend = attention_engaged,
attention_request = attention_backend,
transformer_cache = cache_engaged,
cache_auto = cache_may_toggle,
cache_quant_active = cache_quant_active,
cache_threshold = transformer_cache_threshold,
eager_patched = eager_patched,
speed_deferred = speed_deferred,
compile_cache_ctx = compile_ctx,
hf_token = hf_token,
resolved = resolved,
)
state_committed = True
finally:
# Pre-commit failure: roll back the process-wide mutations (symmetric with _unload_locked).
if not state_committed:
restore_backend_flags(backend_flags_before)
compile_cache.restore(compile_ctx)
gguf_compile.uninstall_all() # idempotent
if eager_patched:
uninstall_patches()
uninstall_arch_patches()
# Free the half-built pipe's VRAM (uncommitted _state -> nothing else reclaims it).
clear_gpu_cache()
logger.info(
"diffusion.loaded: repo=%s base=%s device=%s offload=%s tiling=%s reasons=%s",
repo_id,
base,
device,
effective_policy,
effective_tiling,
"; ".join(plan.reasons),
)
return self.status()
def _load_dense_quant_pipeline(
self,
transformer_cls: Any,
pipeline_cls: Any,
base: str,
device: str,
dtype: Any,
hf_token: Optional[str],
target: DiffusionDeviceTarget,
mode: Optional[str],
fast_accum: Optional[bool] = None,
*,
fam: Optional[DiffusionFamily] = None,
prequant_path: Optional[str] = None,
base_local_dir: Optional[str] = None,
allow_dense_fallback: bool = True,
lora_specs: Optional[list[tuple[str, float]]] = None,
text_encoder_quant: Optional[str] = None,
) -> tuple[Any, str]:
"""Build the opt-in fast pipeline and return ``(pipe, engaged_scheme)``.
Two ways to get the quantized transformer, in order:
1. Pre-quantized: if a checkpoint is configured for the chosen scheme (an explicit
``prequant_path`` or the family's hosted repo), load the already-quantized
weights onto the meta device and assign them in -- the dense bf16 never lands on
the GPU, so the load peak is ~half and the download is smaller.
2. Dense + quantise (fallback): load the DENSE bf16 transformer from the base repo,
place it on the device, and torchao-quantise it in place.
``lora_specs`` bakes LoRA adapters into the build: they attach on the DENSE
transformer (peft's post-quant torchao dispatch needs quantizer metadata a manual
quantize_ never has), then quantize_ converts only the frozen base linears (the
``lora_`` side path is excluded by name), then the loader compiles. That forces the
dense path -- the prequant shortcut is skipped -- so a baked-LoRA load pays the dense
peak. Verified on the Studio stack: scale 0 reproduces the quantized base exactly.
Raises if the scheme is unsupported or quantisation fails, so ``load_pipeline``
catches it and falls back to the GGUF build. Quantisation runs ON the device and
BEFORE the loader compiles the repeated block, so the order stays quantize ->
compile -> placement."""
# 1. Pre-quantized checkpoint, when one is configured for the resolved scheme.
scheme = select_transformer_quant_scheme(target, mode, family = getattr(fam, "name", None))
if scheme is None:
# Bail BEFORE the multi-GB dense download: an unsupported scheme (fp8 on Ampere, nvfp4 off Blackwell) would
# materialise the transformer only to fail at quantize, after eviction. load_pipeline catches this and builds the GGUF pipeline.
raise RuntimeError("transformer quant unsupported for this device/scheme")
if fam is not None and not lora_specs:
# A LoRA bake needs the DENSE transformer (adapters attach before quantize_), so skip prequant.
source = resolve_prequant_source(
fam, scheme, path_override = prequant_path, base_repo = base
)
if source is not None:
transformer = load_prequantized_transformer(
transformer_cls,
base,
source,
device = device,
dtype = dtype,
hf_token = hf_token,
scheme = scheme,
# Reject a checkpoint with a different Linear filter so prequant matches runtime-quant.
min_features = DEFAULT_MIN_LINEAR_FEATURES,
# Only enforced when the caller forces fp8 fast-accum; a checkpoint that baked the other choice falls to the dense path instead of using the baked kernels.
fast_accum = fast_accum,
logger = logger,
)
if transformer is not None:
pipe = self._assemble_pipe(
pipeline_cls,
base,
transformer,
dtype,
hf_token,
device,
base_local_dir,
fam = fam,
te_quant_mode = text_encoder_quant,
target = target,
)
return pipe, scheme
# 2. Fallback: materialise the dense bf16 transformer and quantise it on-device.
if not allow_dense_fallback:
# The memory plan only budgeted the prequant-sized build, so the dense bf16 transformer would exceed it after eviction. Raise to the GGUF build.
raise RuntimeError(
"prequant checkpoint unavailable and the dense transformer does not fit resident"
)
transformer = transformer_cls.from_pretrained(
base,
subfolder = "transformer",
torch_dtype = dtype,
token = hf_token,
cache_dir = hub_cache_dir(),
)
pipe = self._assemble_pipe(
pipeline_cls,
base,
transformer,
dtype,
hf_token,
device,
base_local_dir,
fam = fam,
te_quant_mode = text_encoder_quant,
target = target,
)
if lora_specs:
# Bake the adapters BEFORE quantize_: peft injects its wrappers on the dense Linears (post-quant torchao dispatch would TypeError), then
# quantize_ converts only each wrapper frozen base_layer while the "lora_" side path stays high precision.
baked = self._resolve_lora_set(
[(i, w) for (i, w) in lora_specs if w != 0],
family = getattr(fam, "name", None),
hf_token = hf_token,
)
for name, path, _weight in baked:
pipe.load_lora_weights(path, adapter_name = name)
pipe.set_adapters(
[n for (n, _p, _w) in baked],
adapter_weights = [w for (_n, _p, w) in baked],
)
pipe._unsloth_loras = baked
pipe._unsloth_loras_baked = True
logger.info(
"diffusion.lora_bake: %d adapter(s) attached before %s quantize",
len(baked),
scheme,
)
scheme = quantize_transformer(
pipe,
target,
mode = mode,
family = getattr(fam, "name", None),
fast_accum = fast_accum,
logger = logger,
)
if scheme is None:
raise RuntimeError("transformer quant unsupported for this device/scheme")
return pipe, scheme
@staticmethod
def _assemble_pipe(
pipeline_cls: Any,
base: str,
transformer: Any,
dtype: Any,
hf_token: Optional[str],
device: str,
base_local_dir: Optional[str] = None,
fam: Optional[DiffusionFamily] = None,
te_quant_mode: Optional[str] = None,
target: Any = None,
) -> Any:
"""Assemble the diffusers pipeline around ``transformer`` and place it on ``device``
(a no-op for an already-placed pre-quantized transformer; it moves the companions)."""
if getattr(fam, "name", None) == KREA2_FAMILY_NAME:
# krea ships transformers-5.x configs and no top-level tokenizer files, so from_pretrained dies in the tokenizer; assemble per-component (diffusion_krea2.py).
krea_te = None
if target is not None:
krea_te = te_prequant_pipe_kwargs(
fam,
base,
te_quant_mode = te_quant_mode,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
).get("text_encoder")
pipe = load_krea2_pipeline(
base_local_dir or base,
dtype,
hf_token = hf_token,
transformer = transformer,
text_encoder = krea_te,
)
pipe.to(device)
return pipe
pipe_kwargs: dict[str, Any] = {
"torch_dtype": dtype,
"transformer": transformer,
"cache_dir": hub_cache_dir(),
}
if hf_token:
pipe_kwargs["token"] = hf_token
if getattr(fam, "name", None) == HIDREAM_FAMILY_NAME:
# The repo ships no Llama text_encoder_4; assemble it from the open mirror, as above.
pipe_kwargs.update(
hidream_te4_kwargs(
dtype,
hf_token,
fam = fam,
te_quant_mode = te_quant_mode,
target = target,
)
)
# Same pre-cast TE injection as the other branches: the dense fast path supplies only the transformer, so the companion TE is the big download.
if target is not None:
pipe_kwargs.update(
te_prequant_pipe_kwargs(
fam,
base,
te_quant_mode = te_quant_mode,
target = target,
dtype = dtype,
hf_token = hf_token,
logger = logger,
)
)
pipe = pipeline_cls.from_pretrained(base_local_dir or base, **pipe_kwargs)
pipe.to(device)
return pipe
def _plan_memory(
self,
target: DiffusionDeviceTarget,
single_file_path: Optional[str],
base: str,
fam: DiffusionFamily,
memory_mode: Optional[str],
cpu_offload: bool,
*,
kind: str = "gguf",
repo_id: Optional[str] = None,
transformer_resident_override_mib: Optional[int] = None,
companion_override_mib: Optional[int] = None,
):
"""Build the memory plan for this load: snapshot free device memory and
estimate the model's resident footprint, then let the planner pick an
offload policy + VAE memory savers. Kept on the backend so the cached base
repo (companion text-encoder / VAE) feeds the size estimate.
The size estimate is per-kind: diffusers keeps GGUF weights packed (per-matmul
transient dequant), so a GGUF loads near its on-disk size; a safetensors
single-file loads near its on-disk size (it carries its dtype), except an fp8
transformer file that gets upcast to bf16 on load (~2x resident); and a full
pipeline is one cached download (transformer + companions), already compressed.
``transformer_resident_override_mib`` replaces the file-size transformer estimate
when the loader is planning for a DIFFERENT artifact than the file on disk (the
dense transformer-quant candidate, whose footprint the auto-policy estimates);
``companion_override_mib`` likewise replaces the cached companion total on that
re-plan, so the base repo's PREFETCHED transformer/ shards -- which land in the
same blob cache _companion_cache_bytes sums -- are not counted as companions on
top of transformer_resident_override_mib (a double-count of the transformer)."""
# Settled (max-over-reads) on cuda: a transient foreign allocation otherwise makes an empty card look full and silently declines the resident/quant fast path.
device_memory = settled_snapshot_device_memory(target)
if kind == "pipeline":
# The whole repo is one cached download, so cached bytes are the resident estimate (bnb-4bit/fp8 stay compressed). A LOCAL path is not cached, so sum its on-disk weights.
local_repo = Path(repo_id).expanduser() if repo_id else None
if local_repo is not None and local_repo.is_dir():
cached = self._local_dir_weight_bytes(local_repo, exclude_transformer = False)
else:
cached = self._cache_bytes(repo_id) if repo_id else 0
cached_mib = int(cached // (1024 * 1024)) if cached else None
model_dense_mib = estimate_safetensors_dense_mib(cached_mib)
# A repo can store weights NARROWER than the loaded dtype: ideogram-4 ships its DiTs as raw float8, so cached
# bytes undershoot the bf16 footprint ~2x and auto would OOM. Plan against the size table bf16 total when it knows this repo.
is_narrow_base = bool(repo_id) and repo_id.strip().lower() == fam.base_repo.lower()
if (
not is_narrow_base
and fam.name == IDEOGRAM4_FAMILY_NAME
and local_repo is not None
and local_repo.is_dir()
):
# A local fp8 mirror never string-matches base_repo, so detect fp8 from the shard headers and reserve the bf16 footprint (a local nf4 mirror stays compressed).
is_narrow_base = ideogram4_repo_is_fp8(repo_id)
if is_narrow_base:
table = family_bf16_components_gb(fam, fam.base_repo)
if table is not None:
# Reserve the bf16 footprint from this network-free constant even without a cache estimate, else model_dense_mib stays None ("size unknown -> resident") and the ~54 GB pipeline OOMs.
table_mib = int(sum(table) * (1000.0**3) / (1024.0 * 1024.0))
model_dense_mib = (
table_mib if model_dense_mib is None else max(model_dense_mib, table_mib)
)
companion_mib = None
else:
if transformer_resident_override_mib is not None:
# Planning for the dense-quant candidate: the auto-policy estimate replaces the file-size derivation; companions stay measured from cache.
transformer_resident = transformer_resident_override_mib
elif kind == "single_file":
# An fp8 checkpoint upcasts to bf16 on load (~2x resident); detect from the basename. Excludes the SDXL (single_file_is_pipeline) case, already bf16.
fp8_upcast = not getattr(fam, "single_file_is_pipeline", False) and (
"fp8" in Path(single_file_path).name.lower() if single_file_path else False
)
transformer_resident = estimate_safetensors_dense_mib(
file_size_mib(single_file_path), fp8_upcast = fp8_upcast
)
else:
transformer_resident = estimate_gguf_resident_mib(file_size_mib(single_file_path))
# Companions (VAE + text encoders) load near on-disk size; sum the base-repo cache, or a LOCAL base on-disk weights (the blob cache is empty for a local path).
if companion_override_mib is not None:
# Re-planning the dense candidate: the prefetched transformer/ shards land in the SAME cache _companion_cache_bytes sums, so use the auto-policy estimate instead of double-counting.
companion_mib = companion_override_mib
else:
companion = self._companion_cache_bytes(base)
companion_mib = int(companion // (1024 * 1024)) if companion else None
model_dense_mib = None
if transformer_resident is not None:
model_dense_mib = transformer_resident + (companion_mib or 0)
# Feed the variant hint so estimate_image_runtime_mib sees distilled markers normalized out of fam.name (distilled needs ~15% less headroom).
variant_hint = " ".join(
p
for p in (
fam.name,
Path(single_file_path).name if single_file_path else "",
repo_id or base or "",
)
if p
)
runtime_headroom = estimate_image_runtime_mib(width = None, height = None, family = variant_hint)
return plan_diffusion_memory(
target = target,
device_memory = device_memory,
model_dense_mib = model_dense_mib,
companion_dense_mib = companion_mib,
runtime_headroom_mib = runtime_headroom,
requested_mode = memory_mode,
explicit_offload = cpu_offload,
)
def _workflow_pipe(self, state: _LoadState, class_name: Optional[str], workflow: str) -> Any:
"""The diffusers pipeline for an image-conditioned ``workflow``, built once and
cached. ``Pipeline.from_pipe`` re-wires the loaded text-to-image pipe's resident
modules (transformer/VAE/text-encoder, incl. any compiled/quantised state) into
the workflow pipeline class, so there is no extra VRAM and no reload. Raises a
clear ValueError when the family does not support the workflow."""
if not class_name:
raise ValueError(
f"{workflow} is not supported for the '{state.family.name}' model family."
)
cached = self._aux_pipes.get(class_name)
if cached is not None:
return cached
import diffusers
# torch_dtype=None is load-bearing: from_pipe otherwise recasts EVERY component to fp32, which hard-crashes the
# dense-quant path (torchao tensor-subclass weights cannot swap_tensors). None reuses resident modules at their loaded dtype.
pipe = getattr(diffusers, class_name).from_pipe(state.pipe, torch_dtype = None)
# Publish to the shared aux cache only if THIS load is still current: from_pipe runs under _generate_lock but
# NOT _lock, so an unload can null _state while it builds and caching would hand a later load stale modules.
with self._lock:
if self._state is state:
self._aux_pipes[class_name] = pipe
return pipe
def _controlnet_pipe(self, state: _LoadState, resolved_cn: Any, cancel: threading.Event) -> Any:
"""Build (once, cached) the family's diffusers ControlNet pipeline around the requested
ControlNet model. The ControlNet model is a small extra module loaded via from_pretrained
and cached by id; the pipeline is assembled with ``Pipeline.from_pipe(base,
controlnet=model)`` -- reusing the resident base modules at their loaded dtype (no reload,
no recast; torch_dtype=None for the same reason as _workflow_pipe). Raises a clear
ValueError when the family declares no ControlNet classes."""
fam = state.family
pipe_cls_name = getattr(fam, "controlnet_pipeline_class", None)
model_cls_name = getattr(fam, "controlnet_model_class", None)
if not pipe_cls_name or not model_cls_name:
raise ValueError(f"ControlNet is not supported for the '{fam.name}' model family.")
import diffusers
cn_model = self._cn_models.get(resolved_cn.id)
if cn_model is None:
if cancel.is_set():
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
# resolve_controlnet accepts a bare owner/name without the base trust gate and from_pretrained would execute a malicious pickle, so run the same Hub malware preflight the chat/export loaders use (a local dir is exempt).
# The preflight fails OPEN, so a remote repo also forces safetensors below, closing the pickle RCE vector when the scan could not run.
remote_cn = not getattr(resolved_cn, "is_local", False)
if remote_cn:
from utils.security import evaluate_file_security
_cn_fs = evaluate_file_security(resolved_cn.path, hf_token = state.hf_token or None)
if _cn_fs.blocked:
raise ValueError(_cn_fs.reason)
# Keep at most one ControlNet resident, else swapping ControlNets accumulates until OOM.
if self._cn_models or self._cn_pipes:
self._cn_models.clear()
self._cn_pipes.clear()
clear_gpu_cache()
import torch
# state.dtype is the display string ("bfloat16"), so pass the real dtype and avoid a float32 load.
cn_dtype = getattr(torch, str(state.dtype).replace("torch.", ""), None)
# Force safetensors for an untrusted remote repo: if the Hub scan failed open above, an embedded pickle would still deserialize. A local dir the user chose is exempt.
cn_from_pretrained_kwargs: dict[str, Any] = {"cache_dir": hub_cache_dir()}
if remote_cn:
cn_from_pretrained_kwargs["use_safetensors"] = True
cn_model = getattr(diffusers, model_cls_name).from_pretrained(
resolved_cn.path,
torch_dtype = cn_dtype,
token = state.hf_token or None, # blank -> anonymous
**cn_from_pretrained_kwargs,
)
if cancel.is_set():
# An unload raced the blocking download; bail BEFORE placement so we do not allocate onto a GPU _unload_locked() just freed.
del cn_model
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
# Placement follows the base offload policy: a resident base places it resident, an offloaded base streams it via group offloading. Best-effort; failure -> resident.
if getattr(state, "offload_policy", OFFLOAD_NONE) != OFFLOAD_NONE and (
_offload_controlnet_module(cn_model, state.device, logger)
):
pass
else:
cn_model = cn_model.to(state.device)
if cancel.is_set():
# An unload raced the download and cleared the caches; caching now would pin it.
del cn_model
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
self._cn_models[resolved_cn.id] = cn_model
key = (pipe_cls_name, resolved_cn.id)
pipe = self._cn_pipes.get(key)
if pipe is None:
pipe = getattr(diffusers, pipe_cls_name).from_pipe(
state.pipe, controlnet = cn_model, torch_dtype = None
)
with self._lock:
# Same race as the model cache: an unload may have cleared _cn_pipes while from_pipe ran.
if cancel.is_set() or self._state is not state:
del pipe
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
self._cn_pipes[key] = pipe
return pipe
@staticmethod
def _align_vae_dtype(pipe: Any, denoiser_attr: str = "transformer") -> None:
"""Cast the VAE to the denoiser's compute dtype before an image-conditioned
call. The img2img/inpaint pipelines VAE-encode the input image at the text-
encoder dtype (bf16), but a prior txt2img DECODE may have left the shared VAE
upcast to fp32 (its ``force_upcast`` path), so the encode would mismatch
(bf16 image vs fp32 VAE). Re-aligning here is safe: our families run bf16 or
fp32 only (the fp16 guard promotes fp16), and a later txt2img decode re-upcasts
as needed. ``denoiser_attr`` is ``pipe.transformer`` for DiT families and
``pipe.unet`` for SDXL. Best-effort; a no-op when already aligned."""
denoiser = getattr(pipe, denoiser_attr, None)
vae = getattr(pipe, "vae", None)
if denoiser is None or vae is None:
return
try:
# Read the dtype from the parameters (a compiled nn.Module may hide .dtype), taking the first FLOATING one (a GGUF transformer leading params are packed uint8).
target_dtype = next(
(p.dtype for p in denoiser.parameters() if p.dtype.is_floating_point),
None,
)
if target_dtype is None:
return
if next(vae.parameters()).dtype != target_dtype:
vae.to(dtype = target_dtype)
except (StopIteration, AttributeError, RuntimeError, TypeError):
pass
@staticmethod
def _resolve_lora_set(
specs: list[tuple[str, float]],
*,
family: Optional[str],
hf_token: Optional[str],
cancel: Optional[threading.Event] = None,
) -> tuple[tuple[str, str, float], ...]:
"""Resolve (id, weight) specs to a ``(name, path, weight)`` tuple set for diffusers.
Shared by the generation-time apply path and the quant load-time bake so both produce
IDENTICAL tuples for the same request (the no-op / weight-only comparisons depend on it).
"""
from core.inference import diffusion_lora
resolved = diffusion_lora.resolve_specs(
specs,
family = family,
hf_token = hf_token,
cancel_event = cancel,
)
# diffusers load_lora_weights takes safetensors only; reject a .gguf adapter as a clean 400.
bad = [r.id for r in resolved if r.fmt != "safetensors"]
if bad:
raise ValueError(
"GGUF LoRA adapters are not supported on the diffusers engine "
f"({', '.join(bad)}); use a .safetensors adapter, or the native engine."
)
# Unique adapter names (diffusers requires distinct names; sanitized stems can collide).
uniq: list[tuple[str, str, float]] = []
seen: set[str] = set()
for r in resolved:
name = r.alias
n = 1
while name in seen:
n += 1
name = f"{r.alias}_{n}"
seen.add(name)
uniq.append((name, r.path, r.weight))
return tuple(uniq)
def _apply_loras(
self, state: Any, loras: Optional[list[tuple[str, float]]], cancel: threading.Event
) -> None:
"""Load + activate requested LoRA adapters on ``state.pipe`` (non-fused), or clear
them when none are requested.
The applied set is recorded on the pipe object, so an unchanged selection is a no-op
and a model swap (a fresh pipe with no marker) resets naturally. Never fuses: fusing
breaks on quantized (bnb-4bit / torchao) transformers and blocks live weight tweaks.
A torchao int8/fp8 pipe carries its adapters from the load-time BAKE (attached before
quantize_ + compile). Its module topology is frozen: weight-only changes go through
set_adapters (value-level, compile-guard safe); adding/removing adapters needs a reload
with the new selection, surfaced as a clean 400 here.
"""
from core.inference import diffusion_lora
pipe = state.pipe
current = getattr(pipe, "_unsloth_loras", ())
specs = [(i, w) for (i, w) in (loras or []) if w != 0]
quant_baked = bool(getattr(pipe, "_unsloth_loras_baked", False))
quant = (state.transformer_quant or "").lower()
if quant in ("int8", "fp8", "nvfp4", "mxfp8"):
self._adjust_baked_loras(state, pipe, specs, current, quant_baked, cancel)
return
if not specs:
if current:
try:
pipe.unload_lora_weights()
except Exception: # noqa: BLE001 -- best-effort clear
pass
pipe._unsloth_loras = ()
return
if not diffusion_lora.supports_lora(
engine = "diffusers",
family = getattr(state.family, "name", None),
model_kind = state.kind,
transformer_quant = state.transformer_quant,
compiled = "compiled" in (getattr(state, "speed_optims", ()) or ()),
):
# Name only routes the user can actually reach: a GPU host never selects the native engine on its own (see diffusion_engine_router), so suggesting it without the override is a dead end.
raise ValueError(
"LoRA is not supported for this model/quantisation on the diffusers engine "
"(GGUF-via-diffusers, or a torch.compile'd Speed=default/max load). Reload with "
"transformer_quant int8 or fp8, which rebuilds the GGUF into a LoRA-capable dense "
"transformer, or use a bf16 / bnb-4bit load at Speed=off/eager. To keep the GGUF "
"weights themselves, run the native engine, which a GPU host selects only when "
"UNSLOTH_DIFFUSION_ENGINE=sd_cpp is set."
)
desired = self._resolve_lora_set(
specs,
family = getattr(state.family, "name", None),
hf_token = state.hf_token,
cancel = cancel,
)
uniq = list(desired)
if desired == current:
return
try:
if current:
pipe.unload_lora_weights()
for name, path, _weight in uniq:
pipe.load_lora_weights(path, adapter_name = name)
pipe.set_adapters(
[name for name, _p, _w in uniq], adapter_weights = [w for _n, _p, w in uniq]
)
except Exception as exc: # noqa: BLE001 -- surface as a clean 400
try:
pipe.unload_lora_weights()
except Exception: # noqa: BLE001
pass
pipe._unsloth_loras = ()
raise ValueError(f"Failed to apply LoRA: {exc}") from exc
pipe._unsloth_loras = desired
def _adjust_baked_loras(
self,
state: Any,
pipe: Any,
specs: list[tuple[str, float]],
current: tuple,
quant_baked: bool,
cancel: threading.Event,
) -> None:
"""Generation-time LoRA handling for a torchao-quantized pipe.
The adapters (if any) were baked at load time, before quantize_ + compile, so the
module topology is immutable here. Allowed without a reload: weight tweaks on the
baked set and disabling everything (scale 0 reproduces the quantized base exactly;
set_adapters is value-level, so torch.compile guards absorb it). Anything that would
change topology (adding adapters to a bake-less load, or a different adapter set)
raises a clean 400 telling the client to reload with the new selection.
"""
if not quant_baked:
if not specs:
return # no adapters baked, none requested
raise ValueError(
"This quantized (int8/fp8) load was built without LoRA adapters. Reload the "
"model with the adapter selection to bake it into the quantized transformer."
)
if not specs:
# Disable every baked adapter: scale 0 reproduces the quantized base exactly.
names = [n for (n, _p, _w) in current]
if any(w != 0 for (_n, _p, w) in current):
pipe.set_adapters(names, adapter_weights = [0.0] * len(names))
pipe._unsloth_loras = tuple((n, p, 0.0) for (n, p, _w) in current)
return
desired = self._resolve_lora_set(
specs,
family = getattr(state.family, "name", None),
hf_token = state.hf_token,
cancel = cancel,
)
if desired == current:
return
if [(n, p) for (n, p, _w) in desired] == [(n, p) for (n, p, _w) in current]:
# Same adapters, new weights: value-level change on the baked topology.
pipe.set_adapters(
[n for (n, _p, _w) in desired],
adapter_weights = [w for (_n, _p, w) in desired],
)
pipe._unsloth_loras = desired
return
raise ValueError(
"The LoRA selection changed, but a quantized (int8/fp8) transformer bakes its "
"adapters at load time. Reload the model with the new adapter selection."
)
@staticmethod
def _reset_step_cache(pipe: Any) -> None:
"""Clear the transformer's stateful step cache (FBCache) before a forward.
diffusers keys FBCache residuals by cache context ("cond"/"uncond") on the
long-lived transformer. The context exit does NOT reset them; the end of a
pipeline ``__call__`` does, via ``maybe_free_model_hooks()`` -- but only when
the call RETURNS. A call that raised (an OOM this generate() backs off from, a
cancelled denoise, a failed prior request) leaves its own batch's residual on
the resident transformer, and the next forward's first step then compares
against it: a tensor-shape mismatch when the resolution/batch changed, or a
stale-cache reuse otherwise. The transformer-level reset entry point is
``_reset_stateful_cache`` in diffusers 0.39 (``reset_stateful_hooks`` lives only
on the HookRegistry, so a getattr for it on the transformer is a silent no-op).
Best-effort: a transformer without the hook (uncached load) is a silent no-op."""
transformer = getattr(pipe, "transformer", None)
reset = getattr(transformer, "_reset_stateful_cache", None) or getattr(
transformer, "reset_stateful_hooks", None
)
if callable(reset):
try:
reset()
except Exception: # noqa: BLE001 — reset is best-effort, never fail a generation
pass
def _engage_deferred_speed(self, state: _LoadState) -> None:
"""Engage the deferred `default` speed profile at the start of the 3rd
generation this session.
The load left the pipe fully eager (bit-identical reference); by the 3rd
image repeated use is established, so pay the one-time compile now: eager
patches + attention auto upgrade + regional compile -- exactly what an
unset-speed GGUF load gets at load time. Runs under _generate_lock (the
caller), so no denoise can race the mutation. The flag is cleared FIRST so
a failure never retries per generation; unload cleans everything up via the
same state fields the load-time path uses (backend flags were snapshotted
at load, before any speed layer could mutate them)."""
object.__setattr__(state, "speed_deferred", False)
from .diffusion_eager_patches import install_compile_safe_patches
from .diffusion_arch_patches import install_arch_patches
target = self._resolve_device_target(state.family)
install_compile_safe_patches()
install_arch_patches()
object.__setattr__(state, "eager_patched", True)
# Re-run the load-time selection with the caller ORIGINAL request: auto still upgrades to cuDNN here (speed_active=True), but an explicit backend is honored verbatim.
attention_engaged = apply_attention_backend(
state.pipe,
select_attention_backend(target, state.attention_request, speed_active = True),
logger = logger,
)
object.__setattr__(state, "attention_backend", attention_engaged)
gguf_transformer = state.kind == "gguf" and state.transformer_quant is None
if compile_eligible(target, is_gguf = gguf_transformer, family = state.family):
compile_ctx = compile_cache.begin(
family = state.family.name,
# U-Net families (SDXL) carry the denoiser as pipe.unet.
transformer = getattr(state.pipe, "transformer", None)
or getattr(state.pipe, "unet", None),
dtype = getattr(target, "dtype", None),
# Same GGUF-vs-dense graph distinction as the load-time begin().
quant = "gguf" if gguf_transformer else state.transformer_quant,
attention_backend = attention_engaged,
compile_kwargs = {
# Mirrors the load-time fullgraph decision: an engaged or still-toggleable step cache, or an offload, graph-breaks.
"fullgraph": state.transformer_cache is None
and not state.cache_auto
and state.offload_policy == OFFLOAD_NONE,
"dynamic": True,
"mode": "default",
},
logger = logger,
)
object.__setattr__(state, "compile_cache_ctx", compile_ctx)
speed_applied = apply_speed_optims(
state.pipe,
target,
is_gguf = gguf_transformer,
family = state.family,
speed_mode = SPEED_DEFAULT,
cache_active = state.transformer_cache is not None or state.cache_auto,
offload_active = state.offload_policy != OFFLOAD_NONE,
logger = logger,
)
object.__setattr__(state, "speed_mode", SPEED_DEFAULT)
object.__setattr__(state, "speed_optims", tuple(k for k, v in speed_applied.items() if v))
entry = (state.resolved or {}).get("speed_mode")
if isinstance(entry, dict):
entry["value"] = SPEED_DEFAULT
entry["reason"] = (
"auto: compiled on the 3rd image this session "
"(repeated use amortises the one-time compile)"
)
att = (state.resolved or {}).get("attention_backend")
if isinstance(att, dict) and att.get("source") == "auto":
att["value"] = attention_engaged or "native"
att["reason"] = (
"cuDNN fused attention upgrade" if attention_engaged else "diffusers default"
)
logger.info(
"diffusion.speed: deferred profile engaged on generation 3 "
"(optims=%s, attention=%s)",
",".join(state.speed_optims) or "none",
attention_engaged or "native",
)
def generate(
self,
*,
prompt: str,
negative_prompt: Optional[str] = None,
width: int = 1024,
height: int = 1024,
# Fallbacks; the route always sends the per-model values the UI seeds.
steps: int = 9,
guidance: float = 0.0,
seed: Optional[int] = None,
batch_size: int = 1,
# Batched multi-image generation (diffusion_batched.py): ``prompts`` renders one image per prompt (txt2img only), ``seeds`` one per seed, ``batch_size`` alone derives base..base+n-1.
# Each image gets its OWN torch.Generator, so same-seed repeats at the same batch shape are bit-identical and a solo regeneration matches up to batch-size kernel numerics.
prompts: Optional[list[str]] = None,
seeds: Optional[list[int]] = None,
# Image-conditioned (base64/data-URL): init alone = img2img, init + mask = inpaint. ``strength`` is the denoise strength (0 = keep source, 1 = full redraw); None = txt2img.
init_image: Optional[str] = None,
mask_image: Optional[str] = None,
strength: Optional[float] = None,
# Upscale (hires fix): factor > 1 with an init image enlarges then re-denoises at low strength.
upscale: Optional[float] = None,
# Reference (FLUX.2): additional reference images beyond init_image (a list). Ignored elsewhere.
reference_images: Optional[list[str]] = None,
# LoRA (id, weight) pairs; loaded non-fused and activated for this generation. None/empty clears.
loras: Optional[list[tuple[str, float]]] = None,
# ControlNet (id, control_image_b64, control_type, strength, guidance_start, guidance_end). None = off.
controlnet: Optional[tuple[str, str, str, float, float, float]] = None,
) -> dict[str, Any]:
import torch
from PIL import Image
# Per-generation cancel Event that unload()/a superseding load set (under _lock) to abort just this denoise. _generate_lock is the only lock the denoise holds.
cancel = threading.Event()
with self._generate_lock:
with self._lock:
# An unload / superseding load signalled the active denoise and is waiting for this
# lock. Python locks are not FIFO, so this request can get in first; refuse instead
# of starting a denoise on a pipeline that is already being torn down.
if self._teardown_waiters:
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
state = self._state
if state is None:
raise RuntimeError(DIFFUSION_NOT_LOADED_MSG)
# Register under _lock so unload()/a load can signal THIS generation.
self._active_generate_cancel = cancel
# Publish an active (step 0) state before the slow pre-denoise setup (deferred compile, LoRA resolution, ControlNet build) so a reload mount probe does not read idle. The per-step callback swaps in its own _GenState.
self._gen = _GenState(total_steps = steps)
try:
# The local `state` ref keeps the pipe alive even if unload() nulls _state. Resolve the per-image (prompt, seed) jobs up front: N prompts, one prompt x N seeds, or the legacy single
# prompt deriving base..base+batch_size-1 (matching the native sd.cpp engine, so a gallery recipe replays any batch member alone). A fresh base seed stays < 2**53 for JSON.
jobs, seed = resolve_batch_jobs(
prompt = prompt,
prompts = prompts,
seed = seed,
seeds = seeds,
batch_size = batch_size,
draw_seed = torch.Generator(device = state.device).seed,
)
# Deferred speed auto: engage the compile profile on the 3rd image, before the LoRA/workflow wiring. Best-effort; a failure stays eager and never retries.
# NOT when a LoRA is requested: a compiled transformer rejects LoRA, which would break every LoRA generation on this load.
lora_requested = any(w != 0 for (_id, w) in (loras or []))
# Also stay eager while a PRIOR generation adapters are attached: _apply_loras runs AFTER this, so compiling would bake the adapter in permanently.
loras_attached = bool(getattr(state.pipe, "_unsloth_loras", ()))
if (
state.speed_deferred
and state.generation_count >= 2
and not lora_requested
and not loras_attached
):
try:
self._engage_deferred_speed(state)
except Exception as exc: # noqa: BLE001 — speed is best-effort
logger.warning(
"diffusion.speed: deferred engagement failed, staying eager: %s",
exc,
)
# Apply/adjust LoRA before picking the workflow pipe; from_pipe pipes share the transformer.
self._apply_loras(state, loras, cancel)
# Select the workflow pipeline: txt2img uses the loaded pipe; img2img/inpaint reuse its modules via from_pipe; an edit model own pipe is already the edit pipeline.
pipe = state.pipe
init_pil = mask_pil = None
control_pil = None
cn_scale = cn_gstart = cn_gend = cn_mode = None
ref_extra: list = []
# Validate dependencies up front: mask/upscale/reference need an input image, and reference needs a supporting family (else the combo silently falls back to txt2img).
if init_image is None:
if mask_image is not None:
raise ValueError("mask_image requires an input image (init_image).")
if upscale is not None and upscale > 1.0:
raise ValueError("upscale requires an input image (init_image).")
if reference_images:
raise ValueError("reference_images require an input image (init_image).")
if reference_images and not getattr(state.family, "reference", False):
raise ValueError(
f"Reference images are not supported for the '{state.family.name}' "
"model family."
)
if getattr(state.family, "edit", False):
# Instruction editing: the loaded pipe IS the edit pipeline and always needs an input image; the prompt is the instruction. No mask, no from_pipe.
if init_image is None:
raise ValueError(
f"{state.family.name} is an image-editing model: provide an input image."
)
if mask_image is not None:
# The edit family has no inpaint pipeline; a mask would be silently dropped.
raise ValueError(
f"{state.family.name} is an image-editing model and does not "
"support masks (mask_image)."
)
workflow = "edit"
init_pil = _decode_b64_image(init_image, mode = "RGB")
elif mask_image is not None and init_image is not None:
workflow = "inpaint"
pipe = self._workflow_pipe(state, state.family.inpaint_pipeline_class, workflow)
init_pil = _decode_b64_image(init_image, mode = "RGB")
mask_pil = _decode_b64_image(mask_image, mode = "L")
elif init_image is not None and upscale is not None and upscale > 1.0:
# Upscale (hires fix): enlarge with Lanczos, then re-run img2img at low strength to add detail without redrawing. Shares the img2img pipeline via from_pipe.
workflow = "upscale"
pipe = self._workflow_pipe(state, state.family.img2img_pipeline_class, workflow)
init_pil = _decode_b64_image(init_image, mode = "RGB")
iw, ih = init_pil.size
# Cap the factor, then the absolute output (longest side 2048) to avoid an OOM-scale latent; round to a multiple of 16 (VAE downsample + patch).
factor = max(1.0, min(float(upscale), 4.0))
tw_f, th_f = iw * factor, ih * factor
max_side = 2048
fit = min(1.0, max_side / max(tw_f, th_f))
tw = max(16, int(round(tw_f * fit / 16.0)) * 16)
th = max(16, int(round(th_f * fit / 16.0)) * 16)
# After the cap, the target must still exceed the input (else upscale shrinks it).
if max(tw, th) <= max(iw, ih):
raise ValueError(
f"Upscale would not enlarge this image: its longest side "
f"({max(iw, ih)}px) already meets the {max_side}px output limit. "
f"Use a smaller source image."
)
init_pil = init_pil.resize((tw, th), Image.LANCZOS)
if strength is None:
strength = 0.35 # hires-fix default: preserve content, add detail
elif getattr(state.family, "reference", False) and init_image is not None:
# FLUX.2 reference conditioning: the loaded pipe takes the reference via `image` and generates at the REQUESTED size (no from_pipe, no strength). After inpaint/upscale so those route first.
workflow = "reference"
init_pil = _decode_b64_image(init_image, mode = "RGB")
# Additional references (FLUX.2 combines a list); capped to bound VRAM.
ref_extra = [
_decode_b64_image(x, mode = "RGB") for x in (reference_images or [])[:3]
]
elif init_image is not None:
workflow = "img2img"
pipe = self._workflow_pipe(state, state.family.img2img_pipeline_class, workflow)
init_pil = _decode_b64_image(init_image, mode = "RGB")
else:
workflow = "txt2img"
# ControlNet (diffusers): txt2img only. Builds the family CN pipeline around resident modules and passes a control map.
if controlnet is not None:
from core.inference import diffusion_controlnet
cn_id, cn_image_b64, cn_type, cn_strength, cn_gs, cn_ge = controlnet
# strength 0 disables CN: skip the whole path so a no-op never pays the download/VRAM.
if cn_strength in (None, 0, 0.0):
controlnet = None
else:
if workflow != "txt2img":
raise ValueError(
"ControlNet currently combines with plain text-to-image only, not "
f"the {workflow} workflow."
)
if not diffusion_controlnet.supports_controlnet(
engine = "diffusers",
family = state.family.name,
has_controlnet_pipeline = bool(
getattr(state.family, "controlnet_pipeline_class", None)
),
model_kind = state.kind,
transformer_quant = state.transformer_quant,
):
raise ValueError(
"ControlNet is not supported for this model/quantisation on the "
"diffusers engine (needs a bf16 or bnb-4bit load of a family with a "
"ControlNet pipeline; not GGUF-via-diffusers or torchao fp8/int8)."
)
# Decode + preprocess the control image FIRST so a bad image 400s before any CN download/build, at the OUTPUT size to align with latents.
src = _decode_b64_image(cn_image_b64, mode = "RGB")
control_pil = diffusion_controlnet.preprocess_control(src, cn_type).resize(
(width, height), Image.LANCZOS
)
try:
resolved_cn = diffusion_controlnet.resolve_controlnet(
cn_id, family = state.family.name
)
except FileNotFoundError as exc:
# An unknown CN id -> 400, not 500 (the route maps ValueError).
raise ValueError(str(exc)) from exc
pipe = self._controlnet_pipe(state, resolved_cn, cancel)
workflow = "controlnet"
cn_scale, cn_gstart, cn_gend = cn_strength, cn_gs, cn_ge
# Flux Union CN selects its head by an integer control_mode; map the type.
cn_mode = diffusion_controlnet.union_control_mode(cn_id, cn_type)
# A prompt LIST batches plain text-to-image only: the image-conditioned and ControlNet workflows take one
# conditioning image per call, and a silent broadcast would pair every prompt with it. Multiple SEEDS of one prompt stay valid everywhere.
if uniform_prompt(jobs) is None and workflow != "txt2img":
raise ValueError(
"A prompts list is supported for plain text-to-image only; the "
f"{workflow} workflow takes one prompt per call (seed lists still work)."
)
# Snap odd-sized inputs (and the mask) to a multiple of 16 for workflows whose OUTPUT size comes from the input image (img2img/inpaint/edit); txt2img/reference use the slider and upscale is already /16.
if init_pil is not None and workflow in ("img2img", "inpaint", "edit"):
# img2img/inpaint take output size from the upload, so bound the longest side to 2048 first (else a phone photo drives an OOM-scale latent). edit resizes internally.
if workflow in ("img2img", "inpaint"):
init_pil = _clamp_max_side(init_pil, 2048)
init_pil = _snap_to_multiple(init_pil, 16)
if mask_pil is not None and mask_pil.size != init_pil.size:
from PIL import Image as _PILImage
mask_pil = mask_pil.resize(init_pil.size, _PILImage.NEAREST)
if init_pil is not None:
# Keep the VAE encode dtype consistent with the input image.
self._align_vae_dtype(pipe, state.family.denoiser_attr)
# Pipelines vary in accepted kwargs, so gate every optional one on the signature.
call_params = inspect.signature(pipe.__call__).parameters
kwargs: dict[str, Any] = {
# prompt / generator / num_images_per_prompt are set per chunk below: each chunk carries its own prompt slice and one torch.Generator PER IMAGE.
"num_inference_steps": steps,
# Most pipelines use "guidance_scale"; Qwen-Image uses "true_cfg_scale".
state.family.cfg_kwarg: guidance,
}
if state.family.name == IDEOGRAM4_FAMILY_NAME:
# Ideogram 4 drives CFG via EITHER a constant guidance_scale OR a per-step guidance_schedule (check_inputs rejects both).
# At the advertised defaults drop the constant so the recommended 48-step taper engages; else null the schedule.
if steps == 48 and abs(float(guidance) - 7.0) < 1e-6:
kwargs.pop(state.family.cfg_kwarg, None)
else:
kwargs["guidance_schedule"] = None
if state.family.name == LUMINA2_FAMILY_NAME and "cfg_trunc_ratio" in call_params:
# Lumina 2 card recipe runs the CFG double-forward over only the first quarter of the trajectory (cfg_trunc_ratio=0.25); the pipeline default (1.0) visibly oversaturates output.
kwargs["cfg_trunc_ratio"] = 0.25
if init_pil is not None:
# Reference passes the whole list (FLUX.2 combines); others take the single image.
kwargs["image"] = [init_pil, *ref_extra] if ref_extra else init_pil
if mask_pil is not None and "mask_image" in call_params:
kwargs["mask_image"] = mask_pil
if strength is not None and "strength" in call_params:
kwargs["strength"] = strength
# width/height: txt2img uses the slider; image-conditioned pipes must use the INPUT IMAGE own size, else the
# latents mismatch. Many such pipes drop them entirely, so pass only when accepted, derived from the image.
if workflow in ("txt2img", "reference", "controlnet"):
# These generate at the REQUESTED size (reference/control image resized to match).
kwargs["width"] = width
kwargs["height"] = height
elif init_pil is not None:
iw, ih = init_pil.size
if "width" in call_params:
kwargs["width"] = iw
if "height" in call_params:
kwargs["height"] = ih
if negative_prompt and "negative_prompt" in call_params:
kwargs["negative_prompt"] = negative_prompt
if workflow == "controlnet" and control_pil is not None:
# CN pipeline takes the control map + scale; guidance start/end bound its step range. Every kwarg is signature-gated so a family that omits one still runs.
if "control_image" in call_params:
kwargs["control_image"] = control_pil
elif "image" in call_params: # some CN pipelines name it "image"
kwargs["image"] = control_pil
if "controlnet_conditioning_scale" in call_params and cn_scale is not None:
kwargs["controlnet_conditioning_scale"] = cn_scale
if "control_guidance_start" in call_params and cn_gstart is not None:
kwargs["control_guidance_start"] = cn_gstart
if "control_guidance_end" in call_params and cn_gend is not None:
kwargs["control_guidance_end"] = cn_gend
# Union CN mode index (Flux); only when accepted and the type maps to a mode.
if "control_mode" in call_params and cn_mode is not None:
kwargs["control_mode"] = cn_mode
# Per-forward chunks: the whole job list in ONE forward by default (the measured sweet spot: batch 32 on 4-step
# models, ~10-22x over serial engines), bounded by an explicit batch_size cap; the OOM backoff below halves a failed chunk.
chunks = chunk_jobs(jobs, batch_size)
gen = _GenState(total_steps = steps * len(chunks))
# Steps completed by FINISHED chunks, so the bar spans the whole multi-chunk call (mutable cell: _on_step closes over it).
steps_done = [0]
def _on_step(pipe, step_index, timestep, callback_kwargs):
# Monotonic: a wall-clock adjustment (NTP) mid-denoise would skew the ETA.
now = time.monotonic()
gen.step = steps_done[0] + step_index + 1
if gen.first_step_at == 0.0:
gen.first_step_at = now
gen.eta_seconds = _estimate_eta(
gen.total_steps, gen.step, gen.first_step_at, now
)
# Preempt a long denoise on unload/superseding load (diffusers checks _interrupt).
if cancel.is_set():
pipe._interrupt = True
return callback_kwargs
if "callback_on_step_end" in call_params:
kwargs["callback_on_step_end"] = _on_step
# Re-check an AUTO cache decision against the ACTUAL step count (28 steps gains FBCache, a few-step turbo drops it); explicit choices never toggle.
if state.cache_auto:
# Key on the EFFECTIVE denoise steps: an img2img/upscale request at strength < 1 only denoises a fraction of `steps`, so fold in `strength` (when applied, else the pipe own default) to keep FBCache off the short trajectory.
strength_applied = effective_request_strength(
strength,
init_pil is not None,
"strength" in call_params,
call_params["strength"].default if "strength" in call_params else None,
)
denoise_steps = effective_denoise_steps(steps, strength_applied)
toggled = maybe_toggle_step_cache(
state.pipe,
steps = denoise_steps,
quant_active = state.cache_quant_active,
threshold = state.cache_threshold,
logger = logger,
)
if toggled != state.transformer_cache:
# _LoadState is frozen; the one deliberate in-place update, tracking the pipe-level toggle so status() is truthful.
object.__setattr__(state, "transformer_cache", toggled)
entry = (state.resolved or {}).get("transformer_cache")
if isinstance(entry, dict):
entry["value"] = toggled or "off"
entry["reason"] = (
f"auto: {denoise_steps}-step generation "
+ ("reaches" if toggled else "is below")
+ f" {FBCACHE_MIN_STEPS}"
)
self._gen = gen
images: list[Any] = []
per_image_seeds: list[int] = []
chunk_shapes: list[int] = []
pending = list(chunks)
while pending:
chunk = pending.pop(0)
chunk_kwargs = dict(kwargs)
shared = uniform_prompt(chunk)
generators = [
torch.Generator(device = state.device).manual_seed(s) for _, s in chunk
]
if len(jobs) == 1:
# Single image: scalar prompt + generator, exactly the pre-batching call shape (the regression harness checks this path bit-identical).
chunk_kwargs["prompt"] = shared
chunk_kwargs["generator"] = generators[0]
chunk_kwargs["num_images_per_prompt"] = 1
elif shared is not None:
# Uniform prompt: encode it ONCE and fan out per-image generators.
chunk_kwargs["prompt"] = shared
chunk_kwargs["generator"] = generators
chunk_kwargs["num_images_per_prompt"] = len(chunk)
else:
# Distinct prompts: one image per prompt in a single forward. The negative prompt must be broadcast to match, else
# a pipeline asserts on the length (Z-Image) or fails in the transformer txt/img concat (Qwen-Image, Krea 2, FLUX true-CFG). Pipes that already expand a scalar are unaffected.
chunk_kwargs["prompt"] = [p for p, _ in chunk]
chunk_kwargs["generator"] = generators
chunk_kwargs["num_images_per_prompt"] = 1
if isinstance(chunk_kwargs.get("negative_prompt"), str):
chunk_kwargs["negative_prompt"] = [
chunk_kwargs["negative_prompt"]
] * len(chunk)
# Start every forward from a clean step cache: diffusers only resets FBCache state at the END of a SUCCESSFUL
# __call__, so a raised call (OOM, cancel) leaves a residual the next differently sized forward trips over.
if state.transformer_cache:
self._reset_step_cache(state.pipe)
try:
# inference_mode is faster than no_grad and numerically identical here.
with torch.inference_mode():
out = pipe(**chunk_kwargs).images
except Exception as exc: # noqa: BLE001 — reraised unless a splittable OOM
if len(chunk) < 2 or not is_oom_error(exc):
raise
# OOM backoff: halve the failed chunk and retry; finished chunks keep their images and per-image seeds keep every retry reproducible.
empty_cache = getattr(getattr(torch, "cuda", None), "empty_cache", None)
if callable(empty_cache):
empty_cache()
first_half, second_half = split_chunk(chunk)
pending[:0] = [first_half, second_half]
gen.total_steps += steps # one extra chunk to run
logger.warning(
"diffusion.generate: batch of %d hit OOM; retrying as %d + %d",
len(chunk),
len(first_half),
len(second_half),
)
continue
# A cancelled denoise returns a partial image; don't persist it, nor run remaining chunks.
if cancel.is_set():
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
images.extend(out)
per_image_seeds.extend(s for _, s in chunk)
chunk_shapes.append(len(chunk))
steps_done[0] += steps
# Keep progress ACTIVE through the post-denoise work below: the route persists the image AFTER this returns, so a mount probe reading idle now would refresh the gallery too early. The outer finally clears _gen on any exit.
# Persist the warm torch.compile bundle after the first compiled generation; a STATIC compile makes new artifacts per (w,h,batch), so register this shape first. Idempotent, best-effort.
try:
# Register the dims the forward ACTUALLY compiled with (image-conditioned workflows run at the input image size; see _compile_shape_dims), and every distinct chunk size, since a static compile makes one artifact per batch size too.
reg_width, reg_height = _compile_shape_dims(workflow, init_pil, width, height)
static_shapes = "compiled" in (
state.speed_optims or ()
) and compiled_shapes_are_static(state.pipe, state.speed_mode)
for chunk_batch in sorted(set(chunk_shapes)):
compile_cache.register_shape(
state.compile_cache_ctx,
(reg_width, reg_height, int(chunk_batch)),
static = static_shapes,
)
compile_cache.save(state.compile_cache_ctx, logger = logger)
except Exception: # noqa: BLE001 — cache persistence is best-effort
pass
# Count the finished generation (drives deferred speed); a batch is one generation.
object.__setattr__(state, "generation_count", state.generation_count + 1)
# Return the PIL images (unencoded); the route embeds recipes and persists them. ``seeds`` records each image OWN seed, so every batch member replays individually.
return {
"images": list(images),
"seed": int(seed),
"seeds": [int(s) for s in per_image_seeds],
"repo_id": state.repo_id,
# The adapters ACTUALLY attached for this generation, so the recipe records a load-time bake too: a quantized load bakes adapters before quantize + compile and the generate request then carries none.
"active_loras": _active_lora_pairs(state.pipe),
# The workflow this generation ACTUALLY ran, so a conditioned image is not presented as a plain Create recipe that would replay as something unrelated.
"workflow": workflow,
}
finally:
# Deregister so a later unload/load can't poke a finished generation (if still ours).
with self._lock:
if self._active_generate_cancel is cancel:
self._active_generate_cancel = None
# Sole clear of the published progress state, on every exit, so it stays active through post-denoise work but a crashed generation never leaves the UI stuck "active".
self._gen = None
def generate_progress(self) -> dict[str, Any]:
"""Live per-step progress for an in-flight generation (lock-free read)."""
gen = self._gen
if gen is None or gen.total_steps <= 0:
return {
"active": False,
"step": 0,
"total_steps": 0,
"fraction": 0.0,
"eta_seconds": None,
}
return {
"active": True,
"step": gen.step,
"total_steps": gen.total_steps,
"fraction": gen.step / gen.total_steps, # step is 1..total, never over 1.0
"eta_seconds": gen.eta_seconds,
}
def unload(self) -> dict[str, Any]:
# Abort an in-flight (lock-free) download so unload/eviction returns promptly.
self._cancel_event.set()
with self._lock:
# Abort an in-flight denoise via ITS cancel event.
if self._active_generate_cancel is not None:
self._active_generate_cancel.set()
# Fence queued generations too: they hold no cancel event yet, so the signal above
# cannot reach them.
self._teardown_waiters += 1
# Cancel any in-flight load (its worker checks this token) and drop the marker.
self._load_token += 1
self._loading = None
# Wait for the signalled denoise to exit BEFORE tearing down: _unload_locked uninstalls process-wide state
# (attention patches, GGUF compile hooks, backend flags, compile cache) the denoise still depends on. Mirrors begin_load.
with self._generate_lock:
with self._lock:
self._unload_locked()
# Teardown is done: _state is None, so the fence has nothing left to protect and
# the next generation gets the plain not-loaded message.
self._teardown_waiters -= 1
return self.status()
def _unload_locked(self) -> None:
state = self._state
if state is None:
return
# Restore the process-wide backend flags this load flipped so the next `off` load is bit-identical; compile_cache.restore + gguf_compile.uninstall_all likewise. All idempotent.
restore_backend_flags(state.backend_flags_before)
compile_cache.restore(state.compile_cache_ctx)
gguf_compile.uninstall_all()
if state.eager_patched:
# Lazy import to keep diffusion.py torch-free to import.
from .diffusion_eager_patches import uninstall_patches
from .diffusion_arch_patches import uninstall_arch_patches
uninstall_patches()
uninstall_arch_patches()
# Deliberately NOT unload_lora_weights(): the whole pipe is dropped below, freeing any adapters with it. Both callers hold _generate_lock, so no denoise is in flight.
# Drop the workflow pipes so they do not pin the freed pipeline modules past unload.
self._aux_pipes.clear()
# Drop any ControlNet models + pipelines so the freed load carries no extra modules.
self._cn_pipes.clear()
self._cn_models.clear()
self._state = None
del state
clear_gpu_cache()
def status(self) -> dict[str, Any]:
state = self._state
if state is None:
return {
"loaded": False,
"repo_id": None,
"family": None,
"base_repo": None,
"device": None,
"dtype": None,
"model_kind": None,
"cpu_offload": False,
"offload_policy": None,
"vae_tiling": False,
"memory_mode": None,
"speed_mode": None,
"speed_optims": [],
"text_encoder_quant": None,
"transformer_quant": None,
"attention_backend": None,
"transformer_cache": None,
"workflows": [],
"supports_lora": False,
"supports_controlnet": False,
"resolved": None,
}
from core.inference import diffusion_controlnet, diffusion_lora
return {
"loaded": True,
"repo_id": state.repo_id,
"family": state.family.name,
"base_repo": state.base_repo,
"device": state.device,
"dtype": state.dtype,
"model_kind": state.kind,
"cpu_offload": state.cpu_offload,
"offload_policy": state.offload_policy,
"vae_tiling": state.vae_tiling,
"memory_mode": state.memory_mode,
"speed_mode": state.speed_mode,
"speed_optims": list(state.speed_optims),
"text_encoder_quant": state.text_encoder_quant,
"transformer_quant": state.transformer_quant,
"attention_backend": state.attention_backend,
"transformer_cache": state.transformer_cache,
"resolved": state.resolved,
# Workflows the loaded family supports, so the UI can gate its tabs.
"workflows": _family_workflows(state.family),
"supports_lora": diffusion_lora.supports_lora(
engine = "diffusers",
family = state.family.name,
model_kind = state.kind,
transformer_quant = state.transformer_quant,
compiled = "compiled" in (getattr(state, "speed_optims", ()) or ()),
),
"supports_controlnet": diffusion_controlnet.supports_controlnet(
engine = "diffusers",
family = state.family.name,
has_controlnet_pipeline = bool(
getattr(state.family, "controlnet_pipeline_class", None)
),
model_kind = state.kind,
transformer_quant = state.transformer_quant,
),
}
def _family_workflows(fam: DiffusionFamily) -> list[str]:
"""The workflow ids the diffusers engine can run for ``fam`` (drives UI gating)."""
# Instruction-editing families have no txt2img mode, so expose only "edit".
if getattr(fam, "edit", False):
return ["edit"]
workflows = ["txt2img"]
# Reference families (FLUX.2) add reference conditioning via their pipeline's image arg.
if getattr(fam, "reference", False):
workflows.append("reference")
if getattr(fam, "img2img_pipeline_class", None):
# Upscale runs on the img2img pipeline, so available exactly when img2img is.
workflows.append("img2img")
workflows.append("upscale")
if getattr(fam, "inpaint_pipeline_class", None):
workflows.append("inpaint")
# Outpaint reuses the inpaint pipeline with a padded canvas, so needs one that preserves size.
if getattr(fam, "inpaint_preserves_size", True):
workflows.append("outpaint")
return workflows
def _resolve_base_repo(
repo_id: str, base_repo: Optional[str], fam: DiffusionFamily, hf_token: Optional[str]
) -> str:
"""The companion diffusers repo: caller's base, else the GGUF repo's own
``base_model`` tag, else the family fallback. Shared by both load paths so a
direct ``load_pipeline`` call resolves the variant base the same way.
The base loads via ``from_pretrained``, so it must be trusted -- an explicit
base_repo is already gated at ``validate_load_request``, but the ``base_model``
card tag is attacker-controlled metadata on any remote GGUF repo, so a tag that
is not unsloth/allowlisted/local is dropped in favour of the curated family
default (never fed to ``from_pretrained``), closing the pickle-deserialisation
vector the ControlNet path already guards with evaluate_file_security."""
base = (base_repo or "").strip()
if not base:
tag = _hf_base_model(repo_id, hf_token)
if tag and _is_trusted_diffusion_repo(tag):
base = tag
return resolve_base_repo(fam, base)
def _hf_base_model(repo_id: str, hf_token: Optional[str]) -> Optional[str]:
"""The diffusers base repo from a GGUF repo's ``base_model`` tag, or None.
Lets one family entry cover every variant (Turbo/full, schnell/dev, the
2512 Qwen revision). Skipped for local paths; None on any lookup failure.
"""
if Path(repo_id).expanduser().exists():
return None
try:
from huggingface_hub import HfApi
meta = HfApi().model_info(repo_id, token = hf_token).cardData or {}
except Exception: # noqa: BLE001 — best-effort; fall back to the family default
return None
base = meta.get("base_model")
if isinstance(base, list):
base = base[0] if base else None
return base if isinstance(base, str) and base.strip() else None
def _offload_controlnet_module(cn_model: Any, device: str, logger: Any) -> bool:
"""Stream a ControlNet module through ``device`` via diffusers group offloading.
Used when the base model was loaded with an offload policy: forcing the ControlNet
fully resident with ``.to(device)`` would defeat that low-VRAM placement and can OOM.
Group offloading is applied to this single module (it does not touch the base pipe's
existing hooks), so it is isolated and reversible. Returns True on success; on any
failure the caller falls back to a resident placement, so this never blocks a load."""
try:
import torch
from diffusers.hooks import apply_group_offloading
onload = torch.device(device)
apply_group_offloading(
cn_model,
onload_device = onload,
offload_device = torch.device("cpu"),
offload_type = "block_level",
num_blocks_per_group = 1,
use_stream = onload.type == "cuda",
)
return True
except Exception as exc: # noqa: BLE001 — offload is best-effort; resident is the fallback
if logger is not None:
logger.warning("diffusion.controlnet: group offload failed (%s); loading resident", exc)
return False
def _base_file_downloaded(rfilename: str, *, include_transformer: bool = False) -> bool:
"""True for base-repo files ``from_pretrained`` actually fetches.
The transformer is supplied by the GGUF, and repo docs (``assets/``, the
top-level README/PDF/images) are never downloaded — counting them would peg
the progress estimate above what lands on disk, so the bar would sit short of
100% for the whole pipeline-load phase instead of advancing to "finalizing".
``include_transformer`` admits the ``transformer/`` shards for loads where the
dense transformer-quant path will fetch them anyway (see
``_dense_quant_prefetch_needed``)."""
if rfilename.startswith("transformer/"):
return include_transformer
if "/" not in rfilename: # top-level: only the pipeline manifest is fetched
return rfilename == "model_index.json"
return not rfilename.startswith("assets/")
# Weight extensions the base repo need NOT supply when the single file is the whole pipeline (SDXL): from_single_file(config=base) reads only the base structure, not its weights.
_BASE_WEIGHT_EXTS = (
".safetensors",
".bin",
".ckpt",
".pt",
".pth",
".gguf",
".onnx",
".onnx_data",
".msgpack",
".h5",
".pb",
)
def _base_config_file_downloaded(rfilename: str) -> bool:
"""True for base-repo files needed to BUILD a pipeline structure around a whole-pipeline
single file WITHOUT its weights: config / tokenizer / scheduler JSON, but no weight
tensors (the single file supplies those). Used for ``single_file_is_pipeline`` families."""
if not _base_file_downloaded(rfilename):
return False
return not rfilename.lower().endswith(_BASE_WEIGHT_EXTS)
def _pipeline_file_downloaded(rfilename: str) -> bool:
"""True for files a full-pipeline ``from_pretrained`` fetches.
Like ``_base_file_downloaded`` but for the ``pipeline`` kind, where the repo
supplies its OWN transformer weights, so the ``transformer/`` subfolder is kept.
Top-level docs (README/PDF/images) and ``assets/`` are skipped, and so are
artifacts the torch loader never touches -- ONNX / OpenVINO / Flax exports and
dtype-variant twins (``*.fp16.safetensors``: the loader requests the default
variant) -- so an official repo that ships many formats (e.g. SDXL Base) does
not prefetch tens of GB it will not load.
"""
if "/" not in rfilename: # top-level: only the pipeline manifest is fetched
return rfilename == "model_index.json"
lower = rfilename.lower()
if lower.startswith(("assets/", "onnx/", "openvino/")):
return False
name = lower.rsplit("/", 1)[1]
if name.startswith(("openvino_", "flax_")):
return False
if name.endswith((".onnx", ".onnx_data", ".pb", ".msgpack", ".h5", ".ckpt")):
return False
if ".fp16." in name or ".bf16." in name or ".non_ema." in name:
return False
return True
def _progress(
phase: Optional[str],
bytes_downloaded: int = 0,
bytes_total: int = 0,
fraction: float = 0.0,
*,
error: Optional[str] = None,
) -> dict[str, Any]:
return {
"phase": phase,
"bytes_downloaded": bytes_downloaded,
"bytes_total": bytes_total,
"fraction": fraction,
"error": error,
}
_diffusion_backend: Optional[DiffusionBackend] = None
def get_diffusion_backend() -> DiffusionBackend:
global _diffusion_backend
if _diffusion_backend is None:
_diffusion_backend = DiffusionBackend()
return _diffusion_backend