diffusion_compile_cache: auto mode now saves the Mega-cache bundle after the first compiled generation (UNSLOTH_DIFFUSION_COMPILE_CACHE_SAVE=0 opts out), so users get warm restarts without the distributor env; a bundle hit starts clean (no pointless rewrite of the just-loaded artifacts) and explicit mode 1/on keeps the distributor-style re-save. New register_shape + manifest shape coverage: a STATIC compile produces new artifacts per (width, height, batch), so the generate path registers each generation's shape and an uncovered shape re-dirties the context, growing the bundle to cover every shape the session used. Measured (B200, real backend): Qwen-Image deferred gen-3 hitch 29.1 -> 22.2 s warm with bit-identical output (7.9 MB bundle, ~0.5 s save); SDXL gen-3 115.7 -> 24.7 s and a mid-session 768px recompile 65.8 -> 12.6 s (bundle 63.6 -> 98.7 MB after the 768 re-save). diffusion_speed: U-Net denoisers (UNet2DConditionModel; no _repeated_blocks, so the regional compile never reached them) now get a whole-module STATIC torch.compile on the default tier, plus fused QKV projections and a compiled VAE decode. Measured on SDXL (30 steps / 7.0 / 1024px, 4 prompts, LPIPS vs the bit-exact reference): 6.16 -> 3.14 s end to end (1.96x) at LPIPS 0.035, steady state 0.70-0.88 s/image through the real backend. Rejected on measurement: dynamic=True whole-module (366 s compile for 39.3 ms/step vs static's 73 s for 26.9), regional BasicTransformerBlock only (45.0 ms/step; ResNet convs stay eager), max-autotune + inductor flags (25.9 ms/step for a 445 s warmup), channels-last UNet alone (neutral). DiT tiers unchanged: fused QKV measured exactly neutral under the regional compile (Qwen-Image 6.53 vs 6.52 s), so it stays max-only there, and the DiT VAE decode stays eager (a few % of a DiT generation). compiled_shapes_are_static tells the cache layer which loads are per-shape (max tier, U-Net whole-module). diffusion: register each generation's shape with the compile cache before the save, pass pipe.unet to the cache fingerprint when the pipe has no transformer, and correct the transformer_quant resolved reason on dense loads (it claimed a GGUF transformer was loaded on every non-quantized pipeline load). Tests: 333 passing across the related suites (speed 42, compile_cache 27, cache 40, precision 20, backend, base_precision, transformer_quant, memory); ruff clean. Full measurement record: outputs/image_optim_round2_audit.md.
590 lines
28 KiB
Python
590 lines
28 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
|
|
|
|
"""Opt-in speed optimisations for the local diffusion backend.
|
|
|
|
Off by default, so the default render path stays bit-identical to a plain run (the
|
|
property the regression harness checks). When the operator opts in, this applies the
|
|
near-lossless speedups in the order the diffusers guides recommend
|
|
(channels_last + cudnn.benchmark -> compile, with TF32 / fused-QKV under "max"):
|
|
|
|
off - nothing (default; bit-identical reference).
|
|
eager - everything lossless EXCEPT torch.compile: channels_last VAE +
|
|
cudnn.benchmark + the attention backend + the shared eager monkey-patches
|
|
(fused RMSNorm / AdaLayerNorm + per-arch addcmul fusions, see
|
|
diffusion_eager_patches.py / diffusion_arch_patches.py). The fast first-image
|
|
/ casual-use path -- no compile tax to amortise.
|
|
default - LIGHT compile. For a GGUF model: channels_last + cudnn.benchmark +
|
|
torch.compile of ONLY the dequant op chain
|
|
(``torch.compile(dequantize_gguf_tensor, dynamic=True)``) -- the dequant is
|
|
~70-80% of eager GGUF time, so fusing it gives ~1.24-1.64x for a small
|
|
one-time compile (~7.5-10.4s) and ZERO extra VRAM, resolution-invariant
|
|
(the dequant inputs are fixed-shape weights). For a dense (non-GGUF) model
|
|
there is no dequant, so ``default`` falls back to regional torch.compile of
|
|
the denoiser's repeated block; a U-Net denoiser (SDXL) has no repeated-block
|
|
list, so it gets a whole-module STATIC compile instead (1.61x at LPIPS 0.034
|
|
on SDXL, see ``_UNET_WHOLE_COMPILE``).
|
|
max - the FULL torch.compile: regional max-autotune compile of the denoiser's
|
|
repeated block (which fuses the GGUF dequant AND the matmul/norm/elementwise
|
|
in one graph -- ~3.2x on the GGUF Z-Image transformer, PSNR ~36 dB vs eager,
|
|
well above the Q4 noise floor) plus TF32 matmul and fused QKV projections.
|
|
|
|
Tier rationale: ``default`` is the cheap, always-amortising compile (compile just the
|
|
hot GGUF dequant; the block stays eager) so the first image is fast and VRAM is
|
|
untouched; ``max`` pays the larger regional-compile tax for the bigger warm speedup.
|
|
The compiled dequant is deliberately skipped under ``max`` -- the regional block compile
|
|
subsumes the dequant fusion (a separately-compiled dequant would be traced into that
|
|
graph and break it), so ``max`` runs the stock dequant and lets the block compile it. The
|
|
per-family ``supports_torch_compile`` flag and the bf16/CUDA checks gate regional compile.
|
|
|
|
The backend flags this layer flips (TF32, cudnn.benchmark) are PROCESS-WIDE, so
|
|
``snapshot_backend_flags`` / ``restore_backend_flags`` let the caller capture the
|
|
prior values at load and restore them at unload, keeping a later ``off`` load
|
|
bit-identical instead of inheriting a previous ``max`` run's globals. torch is
|
|
imported lazily.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any, Optional
|
|
|
|
from . import diffusion_gguf_compile as gguf_compile
|
|
|
|
SPEED_OFF = "off"
|
|
SPEED_EAGER = "eager"
|
|
SPEED_DEFAULT = "default"
|
|
SPEED_MAX = "max"
|
|
SPEED_MODES = (SPEED_OFF, SPEED_EAGER, SPEED_DEFAULT, SPEED_MAX)
|
|
|
|
|
|
def snapshot_backend_flags() -> Optional[dict]:
|
|
"""Capture the process-wide torch backend flags this layer may mutate, so the
|
|
caller can restore them on unload. None if torch is unavailable. Each flag is read
|
|
defensively so a build/platform missing one (e.g. no cuda.matmul on CPU/MPS) still
|
|
captures the rest -- otherwise a single missing attribute would skip the whole
|
|
snapshot and a real mutated flag would leak."""
|
|
try:
|
|
import torch
|
|
except Exception: # noqa: BLE001 — no torch -> nothing to snapshot/restore
|
|
return None
|
|
state: dict[str, bool] = {}
|
|
matmul = getattr(getattr(torch.backends, "cuda", None), "matmul", None)
|
|
if matmul is not None and hasattr(matmul, "allow_tf32"):
|
|
state["matmul_tf32"] = bool(matmul.allow_tf32)
|
|
if matmul is not None and hasattr(matmul, "allow_fp16_accumulation"):
|
|
state["matmul_fp16_accum"] = bool(matmul.allow_fp16_accumulation)
|
|
cudnn = getattr(torch.backends, "cudnn", None)
|
|
if cudnn is not None:
|
|
if hasattr(cudnn, "allow_tf32"):
|
|
state["cudnn_tf32"] = bool(cudnn.allow_tf32)
|
|
if hasattr(cudnn, "benchmark"):
|
|
state["cudnn_benchmark"] = bool(cudnn.benchmark)
|
|
inductor_cfg = _inductor_config()
|
|
if inductor_cfg is not None and hasattr(inductor_cfg, "emulate_precision_casts"):
|
|
state["inductor_emulate_precision_casts"] = bool(inductor_cfg.emulate_precision_casts)
|
|
return state
|
|
|
|
|
|
def restore_backend_flags(state: Optional[dict]) -> None:
|
|
"""Restore the flags captured by ``snapshot_backend_flags``. No-op on None. Each
|
|
flag is restored independently so one failure can't leave the others leaked."""
|
|
if not state:
|
|
return
|
|
try:
|
|
import torch
|
|
except Exception: # noqa: BLE001 — no torch -> nothing to restore
|
|
return
|
|
|
|
def _set(obj: Any, attr: str, key: str) -> None:
|
|
if obj is not None and key in state and hasattr(obj, attr):
|
|
try:
|
|
setattr(obj, attr, state[key])
|
|
except Exception: # noqa: BLE001 — best-effort per-flag restore
|
|
pass
|
|
|
|
matmul = getattr(getattr(torch.backends, "cuda", None), "matmul", None)
|
|
_set(matmul, "allow_tf32", "matmul_tf32")
|
|
_set(matmul, "allow_fp16_accumulation", "matmul_fp16_accum")
|
|
cudnn = getattr(torch.backends, "cudnn", None)
|
|
_set(cudnn, "allow_tf32", "cudnn_tf32")
|
|
_set(cudnn, "benchmark", "cudnn_benchmark")
|
|
_set(_inductor_config(), "emulate_precision_casts", "inductor_emulate_precision_casts")
|
|
|
|
|
|
def _inductor_config() -> Any:
|
|
"""``torch._inductor.config`` or None. Resolved as attributes off the imported torch
|
|
module (real torch exposes ``_inductor`` directly after ``import torch``) rather
|
|
than a submodule import, so a stubbed/partial torch (tests, exotic builds) cleanly
|
|
reports None instead of picking a stale real module out of ``sys.modules``."""
|
|
try:
|
|
import torch
|
|
return getattr(getattr(torch, "_inductor", None), "config", None)
|
|
except Exception: # noqa: BLE001 — no inductor -> nothing to snapshot/set
|
|
return None
|
|
|
|
|
|
def normalize_speed_mode(value: Optional[str]) -> str:
|
|
"""Lower/strip a requested speed mode (dashes ok); None / "" -> off."""
|
|
if value is None:
|
|
return SPEED_OFF
|
|
normalized = str(value).strip().lower().replace("-", "_")
|
|
if not normalized:
|
|
return SPEED_OFF
|
|
if normalized not in SPEED_MODES:
|
|
raise ValueError(
|
|
f"Unsupported diffusion speed_mode '{value}'. Use one of: {', '.join(SPEED_MODES)}."
|
|
)
|
|
return normalized
|
|
|
|
|
|
def resolve_speed_mode(
|
|
value: Optional[str],
|
|
*,
|
|
is_gguf: bool,
|
|
dense_default: str = SPEED_OFF,
|
|
) -> str:
|
|
"""The effective speed mode when the caller leaves it UNSET (``None``).
|
|
|
|
A GGUF model defaults to ``default``: it compiles only the hot dequant op chain
|
|
(~70-80% of eager GGUF time) for ~1.24-1.64x at a small one-time compile and zero
|
|
extra VRAM -- a cheap, always-amortising win whose numeric perturbation sits well
|
|
below the quantisation noise floor (the dequant graph is unchanged, just
|
|
Inductor-fused). A dense (non-GGUF) model resolves to ``dense_default``: the image
|
|
backend keeps ``off`` at load (bit-identical first generations, with its own
|
|
deferred engagement after repeated use), while the video backend passes ``default``
|
|
-- a clip denoise runs long enough that the one-time compile always amortises
|
|
within a single generation. An explicit value -- including ``"off"`` -- is always
|
|
honored verbatim."""
|
|
if value is None:
|
|
return SPEED_DEFAULT if is_gguf else dense_default
|
|
return normalize_speed_mode(value)
|
|
|
|
|
|
def compile_eligible(target: Any, *, is_gguf: bool, family: Any) -> bool:
|
|
"""Whether the denoiser's repeated block should be regionally compiled.
|
|
|
|
Only on CUDA (incl. ROCm via supports_default_torch_compile), for a bf16
|
|
transformer, on a compile-friendly family. ``is_gguf`` no longer disqualifies:
|
|
``compile_repeated_blocks`` runs fine on the GGUF transformer (the per-op
|
|
dequant stays eager, the rest of the block compiles) and is ~2.3x faster, so it
|
|
is kept only for signature/logging compatibility."""
|
|
del is_gguf # GGUF is compile-eligible now; param kept for call-site compat.
|
|
if not bool(getattr(target, "supports_default_torch_compile", False)):
|
|
return False
|
|
if not bool(getattr(family, "supports_torch_compile", True)):
|
|
return False
|
|
return _is_bfloat16(getattr(target, "dtype", None))
|
|
|
|
|
|
def _is_bfloat16(dtype: Any) -> bool:
|
|
try:
|
|
import torch
|
|
return dtype is torch.bfloat16
|
|
except Exception:
|
|
return str(dtype).endswith("bfloat16")
|
|
|
|
|
|
def apply_speed_optims(
|
|
pipe: Any,
|
|
target: Any,
|
|
*,
|
|
is_gguf: bool,
|
|
family: Any,
|
|
speed_mode: str = SPEED_OFF,
|
|
cache_active: bool = False,
|
|
offload_active: bool = False,
|
|
logger: Any = None,
|
|
) -> dict[str, bool]:
|
|
"""Apply the opt-in speed optimisations for ``speed_mode`` to a built pipeline,
|
|
BEFORE placement / offload. Returns which optimisations actually engaged. Every
|
|
step is best-effort: a pipeline that doesn't support one is simply skipped.
|
|
|
|
``offload_active`` is the planned offload policy != none: group/model/sequential
|
|
offloading installs ``@torch.compiler.disable``d onload hooks, so the compile must
|
|
drop ``fullgraph`` (same reason as an active step cache) or it crashes at the first
|
|
denoise step."""
|
|
applied = {
|
|
"channels_last": False,
|
|
"cudnn_benchmark": False,
|
|
"tf32": False,
|
|
"fp16_accum": False,
|
|
"fused_qkv": False,
|
|
"compiled": False,
|
|
"compiled_dequant": False,
|
|
"compiled_vae_decode": False,
|
|
}
|
|
mode = normalize_speed_mode(speed_mode)
|
|
# TF32 and cudnn.benchmark are the process-global flags this may flip (TF32 on max,
|
|
# cudnn.benchmark on any non-off CUDA load). The caller snapshots them before this
|
|
# call and restores on unload / failed load via snapshot_backend_flags /
|
|
# restore_backend_flags, so a later `off` load -- or chat inference in the same
|
|
# process -- never inherits them. We keep no separate bookkeeping here.
|
|
if mode == SPEED_OFF:
|
|
return applied
|
|
|
|
on_cuda = getattr(target, "device", None) == "cuda"
|
|
family_allows_compile = bool(getattr(family, "supports_torch_compile", True))
|
|
|
|
# Lossless: a channels-last VAE speeds up its convolutions with no numeric change.
|
|
applied["channels_last"] = _vae_channels_last(pipe, logger)
|
|
|
|
# Near-lossless: let cuDNN autotune the fixed-shape VAE convs (CUDA only). It may
|
|
# pick a different conv algorithm, so it is a "default"-tier (not bit-identical) win.
|
|
if on_cuda:
|
|
applied["cudnn_benchmark"] = _enable_cudnn_benchmark(logger)
|
|
|
|
# Consumer-only: fp16 GEMMs accumulate in fp16 (~2x on GeForce-class parts, whose
|
|
# fp32-accumulate rate is halved; datacenter HBM parts gain nothing and keep the
|
|
# safer fp32 accumulate). Only affects fp16 matmuls -- bf16 loads were measured
|
|
# bit-identical with the flag on across every family (36/36 same-seed A/B cases),
|
|
# so on the quality-neutral tiers the flag engages only when the compute dtype is
|
|
# NOT fp16. On an fp16 pipeline (the pre-Ampere fallback dtype) the same harness
|
|
# measured real same-seed drift (mean 2-5% on SDXL / FLUX), so fp16 compute gets
|
|
# the 2x accumulate only under ``max``, the tier that already trades exactness for
|
|
# speed. Guarded by a per-family deny-list fed by the overflow validation harness
|
|
# and the UNSLOTH_DISABLE_FP16_ACCUM kill switch.
|
|
if on_cuda:
|
|
applied["fp16_accum"] = _enable_fp16_accumulation(
|
|
family, logger, dtype = getattr(target, "dtype", None), speed_mode = mode
|
|
)
|
|
|
|
# --- the compile lever, remapped per tier ----------------------------------------
|
|
# default = LIGHT compile: for a GGUF model, compile ONLY the dequant op chain
|
|
# (~70-80% of eager GGUF time) -- cheap, VRAM-free, resolution-invariant; the
|
|
# transformer block stays eager. A dense model has no dequant, so default falls
|
|
# back to the regional block compile (its only compile lever).
|
|
# max = FULL compile: regional max-autotune compile of the repeated denoiser block
|
|
# (fuses dequant + matmul + norm + elementwise in one graph). It subsumes the
|
|
# dequant fusion, so we do NOT also install the standalone compiled dequant here.
|
|
# eager = no compile at all.
|
|
if mode == SPEED_DEFAULT:
|
|
if is_gguf and on_cuda and family_allows_compile:
|
|
applied["compiled_dequant"] = gguf_compile.install_compiled_dequant(logger)
|
|
elif compile_eligible(target, is_gguf = is_gguf, family = family):
|
|
# A U-Net denoiser (SDXL) fuses QKV BEFORE its whole-module compile:
|
|
# measured 36.3 vs 39.3 ms/step on SDXL under the compile (LPIPS unchanged
|
|
# at 0.033). DiTs measured exactly neutral under the regional compile
|
|
# (Qwen-Image 6.53 vs 6.52 s), so they keep the fuse on the max tier only.
|
|
if _denoiser_unet(pipe) is not None:
|
|
applied["fused_qkv"] = _fuse_qkv(pipe, logger)
|
|
applied["compiled"] = _compile_repeated_blocks(
|
|
pipe,
|
|
logger,
|
|
max_autotune = False,
|
|
cache_active = cache_active,
|
|
offload_active = offload_active,
|
|
)
|
|
elif mode == SPEED_MAX and compile_eligible(target, is_gguf = is_gguf, family = family):
|
|
applied["compiled"] = _compile_repeated_blocks(
|
|
pipe,
|
|
logger,
|
|
max_autotune = True,
|
|
cache_active = cache_active,
|
|
offload_active = offload_active,
|
|
)
|
|
|
|
# A compiled U-Net family also compiles the VAE decode: at SDXL's fast step rate the
|
|
# decode is a real share of each image (measured 4.98 -> 4.25 s over 4 images, LPIPS
|
|
# unchanged). DiT families skip it (the decode is a few % of their generation).
|
|
# dynamic=True keeps it resolution-robust; fullgraph=False tolerates offload hooks.
|
|
if applied["compiled"] and _denoiser_unet(pipe) is not None:
|
|
applied["compiled_vae_decode"] = _compile_vae_decode(pipe, logger)
|
|
|
|
if mode == SPEED_MAX:
|
|
# Near-lossless: TF32 matmul (CUDA only) trades a few mantissa bits for speed.
|
|
if on_cuda:
|
|
applied["tf32"] = _enable_tf32(logger)
|
|
applied["fused_qkv"] = _fuse_qkv(pipe, logger)
|
|
|
|
return applied
|
|
|
|
|
|
def _vae_channels_last(pipe: Any, logger: Any) -> bool:
|
|
vae = getattr(pipe, "vae", None)
|
|
if vae is None or not hasattr(vae, "to"):
|
|
return False
|
|
try:
|
|
import torch
|
|
vae.to(memory_format = torch.channels_last)
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "channels_last", exc)
|
|
return False
|
|
|
|
|
|
# U-Net denoisers ship no ``_repeated_blocks`` (their block mix is heterogeneous), so the
|
|
# regional compile below cannot reach them; these classes instead get a WHOLE-module
|
|
# ``torch.compile`` with STATIC shapes, the one flavor measured worth its warmup. On SDXL
|
|
# (B200, 30 steps / 7.0 / 1024px, 4 prompts; scripts/image_speedmem_bench.py levers +
|
|
# probe): static whole-UNet compile runs 26.9 ms/step vs the 45.9 ms/step bit-exact
|
|
# reference -- **1.61x end to end (6.16 -> 3.83 s) at LPIPS 0.034** -- while dynamic=True
|
|
# compiles 5x slower (366 s vs 73 s cold) for less win (39.3 ms/step), and a regional
|
|
# BasicTransformerBlock compile only reaches 45.0 ms/step (the ResNet convs stay eager).
|
|
# Static shapes mean a recompile per new (height, width, batch); the Mega-cache bundle
|
|
# (diffusion_compile_cache) carries each compiled shape across restarts.
|
|
_UNET_WHOLE_COMPILE: frozenset[str] = frozenset({"UNet2DConditionModel"})
|
|
|
|
|
|
def _denoiser_unet(pipe: Any) -> Any:
|
|
"""The pipe's U-Net denoiser when its class is on the whole-compile list, else None."""
|
|
unet = getattr(pipe, "unet", None)
|
|
if unet is not None and type(unet).__name__ in _UNET_WHOLE_COMPILE:
|
|
return unet
|
|
return None
|
|
|
|
|
|
def compiled_shapes_are_static(pipe: Any, speed_mode: Optional[str]) -> bool:
|
|
"""Whether this load's compiled denoiser artifacts are per-(width, height, batch).
|
|
|
|
The ``max`` tier compiles the regional blocks with dynamic=False, and the U-Net
|
|
whole-module compile is always static; the ``default`` DiT tier compiles
|
|
dynamic=True (one artifact across shapes). The compile-cache layer keys on this to
|
|
re-save its bundle when a session generates at a shape it has not covered yet."""
|
|
mode = normalize_speed_mode(speed_mode)
|
|
if mode == SPEED_MAX:
|
|
return True
|
|
return mode == SPEED_DEFAULT and _denoiser_unet(pipe) is not None
|
|
|
|
|
|
def _denoiser_dits(pipe: Any) -> list:
|
|
"""Every DiT the denoise loop runs each step: the primary ``transformer`` plus a second
|
|
expert some families carry (Ideogram's ``unconditional_transformer`` for its dual-branch
|
|
CFG, an MoE ``transformer_2``). Speed / attention optims must reach ALL of them -- mirroring
|
|
the offload path (diffusion_memory streams the same set) -- else the second DiT runs
|
|
eager / native for every generation while status over-reports the optimisation as engaged."""
|
|
dits: list = []
|
|
for attr in ("transformer", "transformer_2", "unconditional_transformer"):
|
|
m = getattr(pipe, attr, None)
|
|
if m is not None and m not in dits:
|
|
dits.append(m)
|
|
return dits
|
|
|
|
|
|
def _compile_repeated_blocks(
|
|
pipe: Any,
|
|
logger: Any,
|
|
*,
|
|
max_autotune: bool = False,
|
|
cache_active: bool = False,
|
|
offload_active: bool = False,
|
|
) -> bool:
|
|
dits = [
|
|
t for t in _denoiser_dits(pipe) if callable(getattr(t, "compile_repeated_blocks", None))
|
|
]
|
|
unet = _denoiser_unet(pipe) if not dits else None
|
|
if not dits and unet is None:
|
|
return False
|
|
# default: mode="default" + dynamic=True -- fast cold start, robust to resolution
|
|
# changes (no recompile). max: mode="max-autotune-no-cudagraphs" + dynamic=False --
|
|
# Triton autotuning for a few % more on GEMM/conv-heavy models, at a much longer
|
|
# compile and a recompile per new resolution. The CUDA-graph modes (reduce-overhead
|
|
# / max-autotune) are deliberately NOT used: they crash on the regionally-compiled
|
|
# block because its static output buffer is overwritten across denoise steps.
|
|
#
|
|
# fullgraph drops to False when a step cache OR CPU offloading is engaged: both insert
|
|
# an ``@torch.compiler.disable``d function into the forward -- FBCache's per-step
|
|
# decision, and group/model/sequential offload's ``ModuleGroup.onload_`` streaming hook
|
|
# -- i.e. a graph break, which fullgraph=True rejects ("Skip inlining
|
|
# torch.compiler.disable()d function"). The break is cheap and the rest of the block
|
|
# still compiles.
|
|
kwargs: dict[str, Any] = {
|
|
"fullgraph": not (cache_active or offload_active),
|
|
"dynamic": not max_autotune,
|
|
}
|
|
if max_autotune:
|
|
kwargs["mode"] = "max-autotune-no-cudagraphs"
|
|
try:
|
|
import torch
|
|
|
|
# Heterogeneous-block DiTs (e.g. Z-Image) compile ~one graph per distinct block
|
|
# shape through compile_repeated_blocks; Z-Image needs ~11, above dynamo's default
|
|
# recompile_limit of 8. Once the limit is hit a resident load hard-errors under
|
|
# fullgraph (and an offload/cache load silently drops the overflow blocks to eager),
|
|
# so raise it well past that (64) for headroom on larger heterogeneous DiTs. This is
|
|
# diffusers' own documented fix for regional-compile recompilation (their guide bumps
|
|
# cache_size_limit). Deliberately NOT force_parameter_static_shapes=False: it doesn't
|
|
# cut the variant count here and makes each compile ~6x slower (24s -> 143s cold).
|
|
dynamo_cfg = getattr(getattr(torch, "_dynamo", None), "config", None)
|
|
if dynamo_cfg is not None:
|
|
for _limit_attr in ("recompile_limit", "cache_size_limit"): # name varies by torch ver
|
|
if hasattr(dynamo_cfg, _limit_attr):
|
|
setattr(dynamo_cfg, _limit_attr, max(getattr(dynamo_cfg, _limit_attr) or 0, 64))
|
|
# Match eager's intermediate rounding inside inductor's fused pointwise kernels:
|
|
# by default they keep chains in fp32 where eager materialises bf16 between ops,
|
|
# a per-forward rounding delta that a multi-step denoise amplifies chaotically.
|
|
# Measured (B200, scripts/image_speedmem_bench.py, pairwise LPIPS of the
|
|
# compiled tier vs the same-stack eager tier): Qwen-Image 0.019 -> 0.006 at
|
|
# identical speed, FLUX.1-dev 0.046 -> 0.029 at +2% step time, FLUX.2-klein
|
|
# 0.018 -> 0.017 at identical speed; on the video DiT (HunyuanVideo-1.5-720p)
|
|
# full-clip LPIPS vs bit-exact drops 0.221 -> 0.052 at zero cost. Process-
|
|
# global, so snapshot_backend_flags carries it and unload restores the prior
|
|
# value.
|
|
inductor_cfg = _inductor_config()
|
|
if inductor_cfg is not None and hasattr(inductor_cfg, "emulate_precision_casts"):
|
|
inductor_cfg.emulate_precision_casts = True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "compile_repeated_blocks", exc)
|
|
return False
|
|
if unet is not None:
|
|
# Whole-module static compile for the U-Net classes above. fullgraph mirrors the
|
|
# regional decision (an active cache or offload hook graph-breaks, though U-Net
|
|
# pipelines have no CacheMixin so in practice only offload lowers it); dynamic is
|
|
# ALWAYS False -- the measured recipe -- so each new (height, width, batch) pays
|
|
# its own compile, carried across restarts by the Mega-cache bundle.
|
|
# ``Module.compile`` keeps the module identity (in-place ``_compiled_call_impl``),
|
|
# so unload/status/LoRA gating see the same object the eager path had.
|
|
unet_kwargs: dict[str, Any] = {"fullgraph": kwargs["fullgraph"], "dynamic": False}
|
|
if max_autotune:
|
|
unet_kwargs["mode"] = "max-autotune-no-cudagraphs"
|
|
try:
|
|
unet.compile(**unet_kwargs)
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "unet whole-module compile", exc)
|
|
return False
|
|
# Compile every denoiser DiT (a dual-DiT family such as Ideogram runs both each step); a
|
|
# per-DiT failure degrades that one to eager without dropping the others.
|
|
engaged = False
|
|
for transformer in dits:
|
|
try:
|
|
transformer.compile_repeated_blocks(**kwargs)
|
|
engaged = True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "compile_repeated_blocks", exc)
|
|
continue
|
|
# A step cache engaged BEFORE this compile (the production load order) has
|
|
# already wrapped each block's forward in a @torch.compiler.disable'd hook, so
|
|
# the compute branch would run eager on every non-skipped step and forfeit the
|
|
# regional compile entirely. Re-point the hooks' inner forward at compiled
|
|
# wrappers; no-op when no cache hooks are installed. The toggle path (cache
|
|
# engaged after load) is armed by apply_step_cache instead. Lazy import:
|
|
# diffusion_cache imports nothing from this module, but keep the dependency
|
|
# one-directional at import time.
|
|
try:
|
|
from .diffusion_cache import _compile_hooked_block_inners
|
|
_compile_hooked_block_inners(transformer, logger)
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "cache-hook inner compile", exc)
|
|
return engaged
|
|
|
|
|
|
def _compile_vae_decode(pipe: Any, logger: Any) -> bool:
|
|
"""torch.compile the VAE ``decode`` bound method in place (U-Net families only; the
|
|
caller gates). Instance-level assignment: the pipe owns it, unload drops it with the
|
|
pipe, and the module object itself is untouched."""
|
|
vae = getattr(pipe, "vae", None)
|
|
decode = getattr(vae, "decode", None) if vae is not None else None
|
|
if not callable(decode):
|
|
return False
|
|
try:
|
|
import torch
|
|
vae.decode = torch.compile(decode, fullgraph = False, dynamic = True)
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "vae decode compile", exc)
|
|
return False
|
|
|
|
|
|
def _enable_cudnn_benchmark(logger: Any) -> bool:
|
|
try:
|
|
import torch
|
|
torch.backends.cudnn.benchmark = True
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "cudnn_benchmark", exc)
|
|
return False
|
|
|
|
|
|
def _enable_tf32(logger: Any) -> bool:
|
|
try:
|
|
import torch
|
|
|
|
torch.backends.cuda.matmul.allow_tf32 = True
|
|
torch.backends.cudnn.allow_tf32 = True
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "tf32", exc)
|
|
return False
|
|
|
|
|
|
# Families the overflow validation harness (scripts/fp16_accum_validate.py) found to
|
|
# produce non-finite activations or NEW black frames under fp16 accumulation. Empty by
|
|
# measurement: across all six families the harness found no overflow anywhere -- bf16
|
|
# loads are bit-identical with the flag on, and fp16 loads stay finite (their same-seed
|
|
# drift is why fp16 compute is additionally gated to the ``max`` tier below).
|
|
_FP16_ACCUM_DENY: frozenset[str] = frozenset()
|
|
|
|
|
|
def _enable_fp16_accumulation(
|
|
family: Any,
|
|
logger: Any,
|
|
*,
|
|
dtype: Any = None,
|
|
speed_mode: Optional[str] = None,
|
|
) -> bool:
|
|
"""Turn on fp16-accumulated fp16 GEMMs for consumer GPUs, where they run ~2x the
|
|
fp32-accumulate rate (datacenter HBM parts are not throughput-nerfed, so they keep
|
|
the safer default). Gated on: the torch build exposing the flag (2.10+), a
|
|
consumer-class device, the family not being deny-listed by the overflow harness,
|
|
the UNSLOTH_DISABLE_FP16_ACCUM kill switch being unset, and -- when the pipeline
|
|
compute dtype IS fp16, the only case where the accumulator width changes results --
|
|
the ``max`` tier (measured same-seed drift: mean 2-5%; bf16 loads are bit-identical
|
|
so they engage on any tier). The caller's snapshot/restore pair returns the
|
|
process-wide flag to its prior value on unload."""
|
|
import os
|
|
|
|
if os.environ.get("UNSLOTH_DISABLE_FP16_ACCUM", "").strip().lower() in (
|
|
"1",
|
|
"true",
|
|
"yes",
|
|
"on",
|
|
):
|
|
return False
|
|
name = str(getattr(family, "name", family or "")).lower()
|
|
if name in _FP16_ACCUM_DENY:
|
|
return False
|
|
if str(dtype).replace("torch.", "") == "float16" and speed_mode != SPEED_MAX:
|
|
return False
|
|
try:
|
|
import torch
|
|
|
|
matmul = torch.backends.cuda.matmul
|
|
if not hasattr(matmul, "allow_fp16_accumulation"):
|
|
return False
|
|
from .diffusion_transformer_quant import _is_consumer_gpu
|
|
|
|
if not _is_consumer_gpu():
|
|
return False
|
|
matmul.allow_fp16_accumulation = True
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "fp16_accum", exc)
|
|
return False
|
|
|
|
|
|
def _fuse_qkv(pipe: Any, logger: Any) -> bool:
|
|
# Prefer the pipe-level fuse (it covers every component the pipe knows about); else fuse each
|
|
# denoiser DiT directly so a dual-DiT family (Ideogram) fuses BOTH experts, not just the first.
|
|
fn = getattr(pipe, "fuse_qkv_projections", None)
|
|
if callable(fn):
|
|
try:
|
|
fn()
|
|
return True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "fuse_qkv_projections", exc)
|
|
return False
|
|
engaged = False
|
|
for transformer in _denoiser_dits(pipe):
|
|
tfn = getattr(transformer, "fuse_qkv_projections", None)
|
|
if callable(tfn):
|
|
try:
|
|
tfn()
|
|
engaged = True
|
|
except Exception as exc: # noqa: BLE001 — optimisation only
|
|
_warn(logger, "fuse_qkv_projections", exc)
|
|
return engaged
|
|
|
|
|
|
def _warn(logger: Any, what: str, exc: Exception) -> None:
|
|
if logger is not None:
|
|
logger.warning("diffusion.speed: %s failed: %s", what, exc)
|