unsloth/studio/backend/core/inference/diffusion.py
Daniel Han b37ac011e7 Host pre-cast fp8 text encoders for four more families
Round 2 of the hosted TE set, each bit-identical to dense-load-then-cast
and gated through the real backend (marker + status fp8 + same-seed LPIPS
vs dense TEs):

- FLUX.1 T5-XXL (text_encoder_2): 9.52 -> 5.90 GB, one artifact for
  schnell/dev/Krea-dev (T5 shards byte-identical across all three,
  verified sha256). 220 tensors, 144 fp8, LPIPS 0.109.
- Lumina Gemma2-2B: fp32 hub store 10.46 -> 3.20 GB (3.3x download cut).
  288 tensors, 182 fp8, LPIPS 0.041.
- Z-Image Qwen3-4B: 8.04 -> 4.41 GB. 399 tensors, 252 fp8, LPIPS 0.112.
  NOT shared with flux.2-klein-4B: klein retrained layer 35's MLP
  (verified tensor diff, maxdiff 0.86), so klein hosts no entry.
- Krea-2 Qwen3-VL-4B: 8.88 -> 4.83 GB. 713 tensors, 460 fp8, LPIPS 0.082.
  The constructor-assembled krea pipeline takes the encoder directly
  (load_krea2_pipeline text_encoder kwarg); the loader remaps 5.x
  rope_parameters and re-ties weights after assign so the rebuilt encoder
  matches the builder's structure.

HunyuanImage 2.1 reuses the Qwen-Image artifact outright: its Qwen2.5-VL
text encoder is byte-identical (every shard sha256, 16,584,414,544 bytes),
recorded in the new component-level base-equivalence table the checkpoint
validator consults. The injection loop now covers text_encoder.._3 so a
family can host several components. Live check: LPIPS 0.123 vs dense.
2026-07-18 10:25:14 +00:00

3343 lines
173 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,
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_gguf_compile as gguf_compile
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 TE_QUANT_AUTO, normalize_te_quant, quantize_text_encoders
from .diffusion_te_prequant import te_prequant_pipe_kwargs
from .diffusion_vae_quant import VAE_QUANT_AUTO, normalize_vae_quant, quantize_vae
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__)
# A load resolves to one "kind", deciding how the transformer + pipeline is built:
# "gguf" -- single-file GGUF transformer dequantised on-device; companions from base repo.
# "single_file" -- single-file *.safetensors transformer (e.g. fp8), no GGUF dequant; companions from base.
# "pipeline" -- full diffusers repo via from_pretrained, re-applying any embedded quant config.
_MODEL_KINDS = frozenset({"gguf", "single_file", "pipeline"})
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 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 LoRA adapter folder (adapter_config.json + adapter_model.safetensors) is not a
# base checkpoint: from_single_file would fail on the adapter weights AFTER the route
# evicted the resident GPU model. Skip it so the pick stays a pipeline load and 400s in
# validation, before the GPU handoff. Also drop a bare adapter_model.safetensors so a
# config-less adapter export is never reinterpreted as the sole checkpoint.
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 (this is the single decode path for every image-conditioned
# workflow). 4096px covers txt2img's 2048 max, upscales, and outpaint canvases; larger 400s.
max_side = 4096
try:
img = Image.open(io.BytesIO(blob))
# Reject an over-limit image from the header BEFORE img.load() decompresses pixels, so a
# crafted small-payload/huge-dimension file can't spike memory first.
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 base repos that may load as a full (non-GGUF) pipeline despite not being under
# unsloth/. Safetensors-only, no pickle/remote code; exact-match lowercased (typo-squat safe).
# Extend deliberately; never add pickled weights or remote code. The SDXL refiner is
# intentionally NOT here (img2img-only; this backend loads every sdxl repo as base txt2img).
_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 behind each
# 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's guidance-distilled FLUX.1-dev finetune: same arch/layout as dev (FluxPipeline,
# CLIP+T5+ae), gated like dev. Detected as the flux.1 family via the "flux.1" token.
"black-forest-labs/flux.1-krea-dev",
"tongyi-mai/z-image-turbo",
"qwen/qwen-image",
"qwen/qwen-image-2512",
"qwen/qwen-image-edit-2511",
# Krea 2: assembled per-component (diffusion_krea2.py). Turbo = inference; Raw = the
# undistilled base to train LoRAs on (train on Raw, run adapters on Turbo).
"krea/krea-2-turbo",
"krea/krea-2-raw",
# Lumina Image 2.0: standard diffusers layout (Gemma2-2B encoder), safetensors-only,
# loads through the generic from_pretrained pipeline path.
"alpha-vllm/lumina-image-2.0",
# HunyuanImage 2.1: the community diffusers mirror (open, tencent-hunyuan-community
# license), safetensors-only, including the diffusers-native guider components.
"hunyuanvideo-community/hunyuanimage-2.1-diffusers",
# HiDream-I1: open MIT-weights repos, all three variants one family. The Llama TE the
# model_index names comes from the unsloth mirror (diffusion_hidream.py), which the
# unsloth/ org prefix already trusts.
"hidream-ai/hidream-i1-full",
"hidream-ai/hidream-i1-dev",
"hidream-ai/hidream-i1-fast",
# Ideogram 4: no bf16 ships. -fp8 stores the two DiTs as raw float8 (the family base);
# the two nf4 repos are identical bnb-4bit exports (both listed so either id loads).
"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
# Resolved memory profile; 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"); surfaced so the UI can gate
# GGUF-only controls (the dense transformer_quant fast path engages only on gguf).
kind: str = "gguf"
# The opt-in speed profile.
speed_mode: str = SPEED_OFF
speed_optims: tuple = ()
# Process-wide torch backend flags (TF32 / cudnn.benchmark) captured before the speed
# layer mutated them; restored on unload so a later `off` load isn't contaminated.
backend_flags_before: Optional[dict] = None
# Text-encoder quant engaged: "fp8" | "nvfp4" | None.
text_encoder_quant: Optional[str] = None
# VAE quant engaged: "fp8" (layerwise storage) | "fp8_dynamic" (torchao conv) | None. When set,
# the VAE holds fp8 tensor subclasses, so the img2img/inpaint _align_vae_dtype re-cast is
# skipped (it would corrupt them).
vae_quant: Optional[str] = None
# Transformer quant actually engaged on the opt-in dense fast path: "int8" | "fp8"
# | "nvfp4" | "mxfp8" | None. None means the default GGUF transformer was loaded.
transformer_quant: Optional[str] = None
# Attention backend engaged via the diffusers dispatcher, or None for default SDPA.
attention_backend: Optional[str] = None
# The caller's ORIGINAL attention request, carried so the deferred-speed engagement
# re-runs the SAME selection instead of forcing the auto cuDNN upgrade over an explicit pin.
attention_request: Optional[str] = None
# Step cache engaged ("fbcache") or None. Opt-in, for many-step models.
transformer_cache: Optional[str] = None
# AUTO on a cache-capable transformer: generate() re-checks the step count and toggles
# FBCache across FBCACHE_MIN_STEPS. An explicit request (off / fbcache) is never toggled.
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 (any non-off tier); uninstalled on unload.
eager_patched: bool = False
# Deferred speed auto: the load stays eager/bit-identical; generate() engages the `default`
# compile profile at the 3rd generation this session. Cleared once engaged (or failed).
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 from the Hub.
hf_token: Optional[str] = None
# Per-control provenance {control: {value, source, reason}} (source auto/explicit), 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; the ETA rate is measured from there so the
# slower first step (warmup) doesn't skew it.
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
class DiffusionBackend:
"""Holds at most one loaded diffusers pipeline. All mutations are serialised."""
def __init__(self) -> None:
# _lock serialises the small state mutations; status()/load_progress()/
# generate_progress() read lock-free so polling never blocks a slow load.
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 every begin_load/unload so a superseded/cancelled worker neither
# commits its pipeline nor stamps progress onto the current load.
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 (or None), set under _lock to abort THAT
# denoise. Per-generation so a cancel can't be lost to a racing generate nor leak.
self._active_generate_cancel: Optional[threading.Event] = None
# Written by the callback, read lock-free by generate_progress().
self._gen: Optional[_GenState] = None
# Image-conditioned workflow pipes (img2img/inpaint) built via from_pipe around the
# loaded pipe (shared modules, no extra VRAM). Keyed by class name; cleared on unload.
self._aux_pipes: dict[str, Any] = {}
# Loaded ControlNet models (id -> module) and their from_pipe pipelines
# ((class, cn_id) -> pipe), reusing resident base 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 forces load_pipeline onto offload, so the dense build
# never runs and never touches the base transformer/ shards. Widening the prefetch
# would only waste a multi-GB pull -- and a disk-full here has NO GGUF fallback
# (unlike an in-load_pipeline dense failure). Mirror those offload gates.
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 actually take the dense path: resolve the SAME
# candidate load_pipeline re-plans against (which also checks the cache has disk room).
# When disk/scheme/support/a prequant rule it out, don't eagerly pull the shards.
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 candidate loads a small checkpoint, not the dense transformer/ shards,
# so widening for it defeats the savings and can disk-full the fallback-less begin_load pull.
if candidate is None or candidate.prequant:
return False
# Capacity gate: on a device that cannot hold even the candidate's post-quant
# resident set, load_pipeline's re-plan is CERTAIN to decline the dense path, so
# widening would fetch the multi-GB base transformer/ shards (47 GB on Qwen-Image)
# only to run the GGUF as-is. Mirror plan_fits_total_capacity's bar against TOTAL
# capacity (not instantaneous free, which a transient tenant could undercount).
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:
# A deliberately-excluded model gets its stated reason, not the generic
# unknown-family message (which reads like a detection gap and invites a
# family_override retry that would fail deeper and less clearly).
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.)"
)
# Families whose single file IS the whole pipeline (SDXL) have no transformer-only
# GGUF path; reject GGUF here, before the route evicts the current model.
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's dual DiTs) has no transformer-only path;
# a single-file/GGUF load would miss its second DiT. Reject here, 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."
)
# A companion base repo also loads via from_pretrained, so it must clear the same
# trust bar (else a GGUF pick could smuggle in an arbitrary remote base). Gate here.
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 (needs model_index.json); reject a
# non-pipeline local base here, before eviction.
_assert_local_base_is_pipeline(base_repo)
# Reject a bad LOCAL pick now so the route never evicts chat for an unloadable request.
# A path-shaped repo_id is meant to be on disk; a bare "org/name" is a remote HF repo.
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 here, before the handoff: gguf needs .gguf,
# single_file must not be handed a .gguf.
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 an actual .safetensors checkpoint (else it evicts
# chat and only fails in the background from_single_file).
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 a GGUF repo, not a pipeline; loading it as a pipeline
# would evict chat then fail on the missing model_index.json. Reject here.
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,
vae_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 already gated at the route's pre-eviction validate; this re-validation
# is a cheap-fail guard for the resolved repo/family and does not re-gate base_repo.
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,
vae_quant = vae_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 on this thread (both network calls) so
# begin_load returns instantly. Only writer of _loading's fields here.
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
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),
# The dense-quant path otherwise pulls the base transformer/ shards inside the
# locked finalize (unpreemptable); when it can run, pull them in the prefetch here.
include_transformer = kind == "gguf"
and self._dense_quant_prefetch_needed(fam, kwargs),
)
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 (uncommitted _state, so nothing else
# reclaims the VRAM). Guarded so a sticky CUDA error can't skip stamping the real 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 base cache; for a full-pipeline load base IS the repo,
# so count it once (else the bar double-counts).
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)
fraction = min(downloaded / expected, 1.0) if expected > 0 else 0.0
return _progress("downloading", downloaded, expected, fraction)
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 _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,
) -> 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).
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."""
from huggingface_hub import HfApi
api = HfApi()
total = 0
base_files: list[str] = []
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)]
# 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
return total, base_files
# Skip the Hub size lookup for a LOCAL gguf path: model_info would raise on a
# filesystem path and (caught below) skip the base lookup, forcing 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)
total += sum(s.size or 0 for s in info.siblings if s.rfilename == gguf_filename)
# 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)
for s in base_info.siblings:
if base_filter(s.rfilename):
base_files.append(s.rfilename)
total += s.size or 0
except Exception as exc: # noqa: BLE001 — estimate is best-effort
logger.warning("diffusion.size_estimate_failed: %s", exc)
return total, base_files
@staticmethod
def _hub_cache_repo_dir(repo_id: str) -> Path:
"""Local HF hub cache dir for ``repo_id``."""
from huggingface_hub import constants
return Path(constants.HF_HUB_CACHE) / 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,
vae_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 (attached on the dense
# transformer before quantize_ + compile). Ignored by every other load kind: bf16 /
# bnb loads 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/whitespace token must degrade to anonymous, not be passed as a credential
# the Hub client can error on. Normalize once for every branch below.
hf_token = hf_token.strip() if isinstance(hf_token, str) else hf_token
hf_token = hf_token or None
# Validate first (cheap, no torch/diffusers) so a bad family fails even in a no-diffusers
# runtime. Re-sanitize the token (direct callers bypass begin_load).
hf_token = (hf_token.strip() if isinstance(hf_token, str) else hf_token) or None
# base_repo is gated at the route before eviction; this re-validation cheap-fails the
# resolved repo/family and does not re-gate an already-validated base_repo.
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 NOW, before this load evicts the previous
# pipeline. Validate-only for transformer_quant (keep 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)
normalize_vae_quant(vae_quant)
# An explicit Speed="off" (bit-exact) load pins the companions dense too: promoting an
# UNSET or "auto" TE/VAE to auto-quant would silently fp8/int8 them and break the request.
# auto is backend-owned, so both UNSET and "auto" go dense under off; only an explicit
# CONCRETE scheme still forces quant under off.
speed_off = speed_mode is not None and str(speed_mode).strip().lower() == SPEED_OFF
# text_encoder_quant tri-state (mirrors transformer_quant): UNSET ("" / None) or "auto" ->
# auto (best accurate TE scheme, else dense); "none"/"off" -> dense; a concrete scheme forces
# it. Default is auto (dense under off).
if text_encoder_quant is None or str(text_encoder_quant).strip().lower() in ("", "auto"):
text_encoder_quant = "off" if speed_off else TE_QUANT_AUTO
# vae_quant tri-state, same contract: auto -> layerwise fp8 where the family qualifies, else
# dense (fp8_dynamic is explicit-only); none/off -> dense. Also dense under Speed="off".
if vae_quant is None or str(vae_quant).strip().lower() in ("", "auto"):
vae_quant = "off" if speed_off else VAE_QUANT_AUTO
# For a full pipeline the repo itself supplies every component, so it is its
# own base; the single-file kinds resolve the companion base diffusers 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 wheel-only pip
# install can run up to 600s, and doing it under the lock would block unload/cancel that
# whole window. Only an explicit backend pulls a package (auto uses cuDNN/native from
# torch). Best-effort; the authoritative resolve + set under the lock is then 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
# Signal an in-flight denoise to abort, then take _generate_lock to WAIT for it to exit
# before allocating the replacement (a load claims VRAM, so it must not overlap a live pipe).
with self._lock:
# Bail BEFORE signalling a cancel if this load was already superseded, else a stale
# worker would abort an unrelated live generation from the CURRENT model.
if _load_token is not None and _load_token != self._load_token:
raise RuntimeError("Diffusion load was cancelled.")
if self._active_generate_cancel is not None:
self._active_generate_cancel.set()
with self._generate_lock:
with self._lock:
# Re-check: 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()
# 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, so free VRAM is the budget).
# Budgets the GGUF file; the dense-quant fast path is preflighted separately 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 (int8
# min, fp8 on datacenter silicon) over GGUF-as-is; "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" load must stay GGUF-as-is: auto-quant would engage
# int8/fp8 + compile and break the bit-exact request. "off" -> None (GGUF-as-is).
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: load the DENSE bf16 transformer and torchao-quantise it
# (int8/fp8/fp4 tensor cores), which beats GGUF's per-matmul dequant on speed AND
# quality at a higher-memory dense load. CUDA + bf16 + resident fit; ANY failure
# falls back to the GGUF build. GGUF kind only (it has the dense bf16 to materialise).
pipe = None
transformer_quant_engaged = None
quant_plan = None
# The GGUF-size `plan` can mis-budget the fast path two ways, so preflight the real
# footprint BEFORE eviction; both branches need the base repo + a resolved scheme.
dense_declined = False
# False when the memory plan only holds a PREQUANT-sized build: if the prequant
# load then fails, the loader must raise to GGUF instead of materialising the
# dense bf16 transformer the plan never budgeted for.
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
# (int8/fp8 ~half bf16; a prequant never materialises dense). Re-plan
# against the candidate's estimate -- 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 the candidate
# for the dense build it will actually 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 auto-policy's companion estimate so the prefetched base
# transformer/ shards in the cache 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; a fresh
# snapshot cannot change that, so don't waste a retry.
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 device capacity with the standard
# reserve + resident margin, yet the instantaneous free reading
# said no: a transient foreign allocation (measured on B200:
# ~100 GB held for under a minute on an idle card) must not
# force the GGUF fallback. Re-snapshot (settled) and replan
# once before declining.
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 already declined resident; a prequant-sized
# replan says nothing about the (larger) dense transformer.
if candidate.prequant:
dense_fallback_allowed = False
else:
# The GGUF fits resident, but this path first materialises the base's dense
# bf16 transformer (bigger), so re-check the fit against THAT -- a card that
# fits the GGUF but not the dense must skip the fast path up front, not OOM
# after eviction. A prequant loads a small file (no dense), so skip the
# re-check there. _dense_transformer_resident_bytes returns 0 if shards are absent.
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 as prequant here, or it skips the dense-fit re-check
# and OOMs materialising the dense transformer after eviction.
prequant = (
# A LoRA bake skips the prequant shortcut (adapters attach on the
# dense transformer), so the dense-fit re-check must gate the fast
# path exactly as if no prequant source 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 entirely (as before); with
# one, the small prequant load proceeds and 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 keeps the
# partially-built dense transformer/pipe alive, blocking VRAM reclaim.
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 ()))
):
# The client asked for adapters BAKED into a quantized build, and that
# build was declined (memory plan) or failed. The GGUF fallback cannot
# carry the adapters, so completing it would return HTTP success while
# silently generating without the requested LoRAs -- wrong output with
# no signal. Fail the load with the recovery options instead.
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
# any embedded quantization_config (e.g. bnb-4bit).
if fam.name == KREA2_FAMILY_NAME:
# krea ships transformers-5.x configs the 4.x line can't parse; assemble
# per-component (see diffusion_krea2.py). The constructor path never
# sees pipe_kwargs, so the pre-cast TE is handed in directly.
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 (see diffusion_ideogram4.py).
pipe = load_ideogram4_pipeline(repo_id, dtype, hf_token = hf_token)
else:
pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype}
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 (diffusion_hidream.py).
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 (when the family ships one and
# the runtime cast would engage) skips the dense TE download; the
# later 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 (its
# sweep re-pulls files the scoped prefetch skipped: 24 GB per FLUX.1).
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, so load it
# through the pipeline class; ``config`` points at the base repo so diffusers
# builds the correct structure around the single-file weights.
sf_pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "config": base}
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,
}
if kind == "gguf":
# Dequantise the GGUF transformer on-device at the compute dtype.
sf_kwargs["quantization_config"] = diffusers.GGUFQuantizationConfig(
compute_dtype = dtype
)
# 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}
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 the full-pipeline branch: the GGUF
# supplies the transformer, so the companion 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` (compile ~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 and
# loses to GGUF), so force at least `default` when quant engaged.
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 (a one-off image shouldn't pay
# the 25-60s compile), but generate() engages `default` on the 3rd image, where
# repeated use amortises it. Only when speed was unset, nothing forced compile, and this device can compile.
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).
# Snapshot the global backend flags (TF32/cudnn.benchmark) 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); when engaged, compile drops
# fullgraph (graph break). Tri-state: unset/"auto" -> step-count policy decides
# (engage when the DEFAULT schedule reaches 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 (a non-CacheMixin one keeps fullgraph).
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). _unload_locked reverses it via
# _state, so a pre-commit failure would leak it; the try/finally below restores on failure.
# gguf_transformer: the GGUF-specific compiled dequant applies only when the GGUF
# was actually loaded. On the dense fast path gguf_filename is still set (fallback)
# but pipe.transformer is dense (needs REGIONAL block compile), so treat it 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 (qwen _modulate / z-image residual, etc.);
# 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 and reused. A miss is silent -> local compile.
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),
quant = 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 cached bundle must key on the same fullgraph 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 couldn't 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) (opt-in), before placement so
# offload moves the smaller weights. Family drives int8's 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,
)
# Quantise the dense VAE (opt-in fp8 layerwise / fp8_dynamic torchao conv),
# before placement so offload moves the smaller weights. Image families never
# force-fp32 their VAE; auto skips torchao under offload.
vae_quant_engaged = quantize_vae(
pipe,
target,
mode = vae_quant,
family = fam.name,
offload_active = plan.offload_policy != OFFLOAD_NONE,
force_fp32 = False,
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",
),
"text_encoder_quant": (
text_encoder_quant,
te_quant or "off",
"dense (no accurate scheme for this GPU / disabled)"
if te_quant is None
else "auto-selected for this GPU + family"
if text_encoder_quant == TE_QUANT_AUTO
else "requested",
),
"vae_quant": (
vae_quant,
vae_quant_engaged or "off",
"dense (no accurate scheme for this GPU / disabled)"
if vae_quant_engaged is None
else "auto-selected for this GPU + family"
if vae_quant == VAE_QUANT_AUTO
else "requested",
),
"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,
vae_quant = vae_quant_engaged,
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 otherwise 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
# the prequant shortcut is skipped when adapters were requested.
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; materialising the dense
# bf16 transformer here 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
)
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 (the post-quant torchao dispatch would TypeError on a manually
# quantized module), then quantize_ converts only each wrapper's 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
# Pipeline.from_pretrained dies in the tokenizer (vocab_file = None); assemble
# per-component like every other krea load path (see 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}
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
# (diffusion_hidream.py) exactly like the full-pipeline load branch.
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 full-pipeline and GGUF 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 at the wrong instant
# otherwise makes an empty card look full and silently declines the resident/quant fast
# path (see settled_snapshot_device_memory).
device_memory = settled_snapshot_device_memory(target)
if kind == "pipeline":
# The whole repo is one cached download; cached bytes are the resident estimate
# (bnb-4bit/fp8 stay compressed). A LOCAL path isn't 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's base ships its
# two DiTs as raw float8, so cached bytes undershoot the bf16 footprint ~2x and auto
# would OOM. When the size table knows the bf16 total for THIS repo, plan against the larger.
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 when the
# cache estimate is absent; else model_dense_mib stays None ("size unknown ->
# resident") and the ~54 GB fp8 pipeline OOMs a card offload would have fit.
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 (not the file on disk): the auto-policy's
# 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's 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's companion estimate
# instead of double-counting the transformer.
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 (basename + repo) so estimate_image_runtime_mib sees distilled
# markers ("turbo"/"schnell") 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 upcasts the reused bf16 modules and hard-crashes the dense-quant path (a
# torchao+compiled transformer's tensor-subclass weights can't swap_tensors). None reuses
# the resident modules at their loaded dtype (the point of from_pipe).
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; caching then
# would hand a wrapper over stale modules to a later load.
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 deserializes it (a malicious pickle would execute), so run the same
# Hub malware preflight the chat/export loaders use. A local dir is exempt (fail-open).
if not getattr(resolved_cn, "is_local", False):
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: evict the previous module + wrapper first,
# or swapping ControlNets within a load 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"), not a torch.dtype; pass the real
# dtype so diffusers loads at the base compute dtype, not float32 (extra VRAM).
cn_dtype = getattr(torch, str(state.dtype).replace("torch.", ""), None)
cn_model = getattr(diffusers, model_cls_name).from_pretrained(
resolved_cn.path,
torch_dtype = cn_dtype,
token = state.hf_token or None, # blank -> anonymous
)
if cancel.is_set():
# An unload raced the blocking download; bail BEFORE placement so we don't
# allocate onto a GPU _unload_locked() just freed.
del cn_model
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
# Placement follows the base's 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; caching now would pin a pipeline over the unloaded base.
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",
vae_quant: Optional[str] = None,
) -> 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.
Skipped when ``vae_quant`` engaged: the fp8 tensor subclasses mishandle
``.to(dtype=...)`` (torchao rejects it), and the VAE already runs at the compute
dtype under fp8, so the re-align is both harmful and unnecessary."""
if vae_quant is not None:
return
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 plain/compiled nn.Module may hide .dtype).
# Take the first FLOATING dtype (a GGUF transformer's 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
# This branch's resolve_specs has no family-scoped catalog; the parameter is kept so
# the caller code stays identical across branches.
del family
resolved = diffusion_lora.resolve_specs(
specs,
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 ()),
):
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). Use a bf16 "
"or bnb-4bit load at Speed=off/eager, or the native engine for GGUF models."
)
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 generation.
diffusers keys FBCache residuals by cache context ("cond"/"uncond") on the
long-lived transformer, and neither the pipeline nor the context exit resets
them. 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), and no pipeline calls it.
This backend reuses one resident pipe across generations, so without a reset the
next generation's first step compares its first-block residual against the
PREVIOUS request's -- a tensor-shape mismatch when the resolution/batch changed,
or a stale-cache reuse otherwise. 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's ORIGINAL request (not a bare
# None): an explicit backend must survive the deferred upgrade. auto still upgrades
# to cuDNN here (speed_active=True), but an explicit "native"/"sage"/"flash" is
# honored verbatim rather than silently replaced by the auto cuDNN choice.
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),
quant = 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,
# 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 (registered under
# _lock below) to abort just this denoise. _generate_lock is the only lock the denoise holds.
cancel = threading.Event()
with self._generate_lock:
with self._lock:
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 now, before the slow pre-denoise setup (deferred
# compile, LoRA resolution, ControlNet build), so a reload's mount probe doesn't read
# idle while this generation holds _generate_lock and let a second generate queue
# behind it. The per-step callback swaps in its own _GenState at denoise start.
self._gen = _GenState(total_steps = steps)
try:
# The local `state` ref keeps the pipe alive even if unload() nulls _state.
generator = torch.Generator(device = state.device)
if seed is None:
# Keep the seed in JS's safe-integer range (< 2**53) so it round-trips
# through JSON and reproduces the image (a raw 64-bit seed loses precision).
seed = generator.seed() & ((1 << 53) - 1)
else:
seed = int(seed)
generator.manual_seed(seed)
# Deferred speed auto: engage the compile profile on the 3rd image, before the LoRA/
# workflow wiring (load-time ordering). Best-effort; a failure stays eager and never retries.
# NOT when a LoRA is requested: a compiled transformer rejects LoRA, so compiling
# here would permanently 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's adapters are attached: _apply_loras runs
# AFTER the engage below, so compiling would bake the adapter in and the swallowed
# unload_lora_weights() would leave it active forever. Defer until a later gen.
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's 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; always needs an
# input image, 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, txt2img's max)
# 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; output size
# from the sliders. After inpaint/upscale so a mask/upscale on a reference family routes right.
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 (not img2img/inpaint/edit). Builds the
# family's 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. Control map 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)
# Snap odd-sized inputs to a multiple of 16 for workflows whose OUTPUT size comes
# from the input image (img2img/inpaint/edit); txt2img/reference use the slider,
# upscale already produced a /16 target. Mask is matched to the snapped image.
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. Pass the engaged
# vae_quant so a quantised (fp8) VAE skips the re-align (must not be re-cast).
self._align_vae_dtype(pipe, state.family.denoiser_attr, state.vae_quant)
# 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": prompt,
"num_inference_steps": steps,
# Most pipelines use "guidance_scale"; Qwen-Image uses "true_cfg_scale".
state.family.cfg_kwarg: guidance,
"generator": generator,
# Whole batch in one forward pass; all share this call's seed.
"num_images_per_prompt": batch_size,
}
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's card recipe runs the CFG double-forward only over the FIRST
# quarter of the trajectory (cfg_trunc_ratio=0.25); the pipeline default (1.0)
# applies it everywhere, visibly oversaturating output. Constant card value.
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's own size (a differing slider mismatches the latents). Many img2img/inpaint
# 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
gen = _GenState(total_steps = steps)
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 = 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 (a 28-step request
# 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 folding in `strength` (only when
# actually applied) keeps FBCache off the short trajectory. When strength is
# omitted the pipe's own default (< 1) still applies, so key on that too.
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}"
)
# Start each generation from a clean step cache: prior FBCache residuals would
# otherwise be compared against this first step (shape mismatch / stale reuse).
if state.transformer_cache:
self._reset_step_cache(state.pipe)
self._gen = gen
try:
# inference_mode is faster than no_grad and numerically identical here.
with torch.inference_mode():
images = pipe(**kwargs).images
finally:
self._gen = None
# A cancelled denoise returns a partial/garbage image; don't persist it.
if cancel.is_set():
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
# 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
# (an uncovered shape re-dirties the context and the save rewrites the bundle).
# Idempotent + best-effort.
try:
# Register the dims the forward ACTUALLY compiled with (image-conditioned
# workflows run at the input image's size, not the slider; see _compile_shape_dims).
reg_width, reg_height = _compile_shape_dims(workflow, init_pil, width, height)
compile_cache.register_shape(
state.compile_cache_ctx,
(reg_width, reg_height, int(batch_size)),
static = "compiled" in (state.speed_optims or ())
and compiled_shapes_are_static(state.pipe, state.speed_mode),
)
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.
return {"images": list(images), "seed": int(seed), "repo_id": state.repo_id}
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
# Drop the published progress state, covering a setup-time error that skips
# the inner finally. Safe under _generate_lock.
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()
# 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 but its pipe ref doesn't pin. The denoise holds _generate_lock
# for its whole body, so acquiring it here blocks until it has finished. Mirrors begin_load.
with self._generate_lock:
with self._lock:
self._unload_locked()
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() here: the whole pipe is dropped below, freeing any
# LoRA adapters with it. Both callers hold _generate_lock across this teardown, so no denoise
# is in flight.
# Drop the workflow pipes so they don't pin the freed pipeline's 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,
"vae_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,
"vae_quant": state.vae_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 file 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 repo's structure
# (config/tokenizer/scheduler) and takes the weights from the single file.
_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