diff --git a/scripts/perf_levers_probe.py b/scripts/perf_levers_probe.py new file mode 100644 index 0000000000..3964ab46ad --- /dev/null +++ b/scripts/perf_levers_probe.py @@ -0,0 +1,238 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Measure the next-phase diffusion levers on the real model, vs today's compiled baseline. + +Variants (Z-Image dense bf16, regional compile = the shipped "default" speed profile): + baseline -- channels_last + compile_repeated_blocks (reference image) + inductor_flags -- + the lossless inductor autotune flags (conv_1x1_as_mm, + coordinate_descent_tuning(+all_dirs), epilogue_fusion=False) + attn_cudnn -- + set_attention_backend("_native_cudnn") (exact) + attn_flash4 -- + set_attention_backend("flash_4_hub") (exact, SM100) + attn_sage -- + set_attention_backend("sage") (INT8 QK, quantized) + fbcache -- + First-Block-Cache (threshold 0.12) (few-step headroom test) + +Reports median latency, vs-baseline speedup, peak VRAM, and LPIPS vs baseline. One CUDA GPU.""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +import numpy as np + +BASE = "Tongyi-MAI/Z-Image-Turbo" +PROMPT = "A cinematic photograph of a red fox in a snowy forest at dawn, highly detailed" +OUT = Path(__file__).resolve().parent.parent / "outputs" / "quant_research" / "perf_levers_images" + + +_LP = {"fn": None} + + +def _lpips(ref, arr): + try: + import lpips + import torch + + # Keep the metric model on CPU: caching it on CUDA leaves it resident across + # variants, and each run resets peak-memory stats, so its VRAM would be charged + # to (and reduce headroom for) every later variant's measurement. + if _LP["fn"] is None: + _LP["fn"] = lpips.LPIPS(net = "alex", verbose = False).eval() + + def t(x): + return torch.from_numpy(x).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0 + + with torch.no_grad(): + return float(_LP["fn"](t(ref), t(arr)).item()) + except Exception as exc: # noqa: BLE001 + print(f" (lpips: {type(exc).__name__})", flush = True) + return None + + +def _set_inductor_flags(): + import torch._inductor.config as ic + + ic.conv_1x1_as_mm = True + ic.coordinate_descent_tuning = True + ic.coordinate_descent_check_all_directions = True + ic.epilogue_fusion = False + try: + ic.force_fuse_int_mm_with_mul = True + except Exception: # noqa: BLE001 + pass + + +def _reset_inductor_flags(): + import torch._inductor.config as ic + + ic.conv_1x1_as_mm = False + ic.coordinate_descent_tuning = False + ic.coordinate_descent_check_all_directions = False + ic.epilogue_fusion = True + # Reset the int-mm fusion flag too, or it leaks from the inductor_flags variant into + # every later compiled row and the attention/fbcache measurements stop being isolated. + try: + ic.force_fuse_int_mm_with_mul = False + except Exception: # noqa: BLE001 + pass + + +def _load(): + import diffusers + import torch + + t = diffusers.ZImageTransformer2DModel.from_pretrained( + BASE, subfolder = "transformer", torch_dtype = torch.bfloat16 + ) + pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype = torch.bfloat16, transformer = t) + pipe.to("cuda") + try: + pipe.vae.to(memory_format = torch.channels_last) + except Exception: # noqa: BLE001 + pass + return pipe + + +def _gen(pipe, steps, seed, res): + import torch + + g = torch.Generator(device = "cuda").manual_seed(seed) + torch.cuda.synchronize() + t0 = time.time() + img = pipe( + prompt = PROMPT, + width = res, + height = res, + num_inference_steps = steps, + guidance_scale = 0.0, + generator = g, + ).images[0] + torch.cuda.synchronize() + return img, time.time() - t0 + + +def _median(xs): + return sorted(xs)[len(xs) // 2] + + +def run( + tag, + steps, + seed, + res, + iters, + *, + attn = None, + fbcache = None, + inductor = False, +): + import torch + + torch.compiler.reset() + torch.cuda.empty_cache() + torch.cuda.reset_peak_memory_stats() + _reset_inductor_flags() + if inductor: + _set_inductor_flags() + pipe = _load() + note = "" + if attn is not None: + try: + pipe.transformer.set_attention_backend(attn) + except Exception as exc: # noqa: BLE001 + note = f"attn({attn})={type(exc).__name__}:{str(exc)[:60]}" + print(f" [{tag}] {note}", flush = True) + del pipe # free the resident pipe so a skipped variant doesn't leak VRAM + torch.cuda.empty_cache() + return None + if fbcache is not None: + try: + from diffusers.hooks import FirstBlockCacheConfig, apply_first_block_cache + apply_first_block_cache(pipe.transformer, FirstBlockCacheConfig(threshold = fbcache)) + except Exception as exc: # noqa: BLE001 + print(f" [{tag}] fbcache={type(exc).__name__}:{str(exc)[:60]}", flush = True) + del pipe + torch.cuda.empty_cache() + return None + try: + pipe.transformer.compile_repeated_blocks(fullgraph = True, dynamic = True) + except Exception as exc: # noqa: BLE001 + print(f" [{tag}] compile={type(exc).__name__}:{str(exc)[:60]}", flush = True) + try: + _gen(pipe, steps, seed, res) # warmup / compile + except Exception as exc: # noqa: BLE001 + import traceback + + traceback.print_exc() + print(f" [{tag}] FAILED first gen: {type(exc).__name__}:{str(exc)[:80]}", flush = True) + del pipe + torch.cuda.empty_cache() + return None + dts, img = [], None + for _ in range(iters): + img, dt = _gen(pipe, steps, seed, res) + dts.append(dt) + peak = torch.cuda.max_memory_allocated() / 1e9 + arr = np.array(img) + OUT.mkdir(parents = True, exist_ok = True) + img.save(OUT / f"{tag}.png") + del pipe + torch.cuda.empty_cache() + return _median(dts), arr, peak + + +def main(argv = None) -> int: + p = argparse.ArgumentParser() + p.add_argument("--steps", type = int, default = 8) + p.add_argument("--res", type = int, default = 1024) + p.add_argument("--seed", type = int, default = 42) + p.add_argument("--iters", type = int, default = 3) + args = p.parse_args(argv) + s, r, seed, it = args.steps, args.res, args.seed, args.iters + + print(f"== perf levers (Z-Image dense, {r}px, {s} steps) ==", flush = True) + base = run("baseline", s, seed, r, it) + if base is None: + print("baseline FAILED", flush = True) + return 1 + bmed, ref, bpeak = base + print(f" baseline {bmed:.3f}s peak={bpeak:.1f}G", flush = True) + rows = [("baseline", bmed, bpeak, 0.0)] + + variants = [ + ("inductor_flags", dict(inductor = True)), + ("attn_cudnn", dict(attn = "_native_cudnn")), + ("attn_flash4", dict(attn = "flash_4_hub")), + ("attn_sage", dict(attn = "sage")), + ("attn_sage_inductor", dict(attn = "sage", inductor = True)), + ("fbcache_0p12", dict(fbcache = 0.12)), + ] + for tag, kw in variants: + out = run(tag, s, seed, r, it, **kw) + if out is None: + rows.append((tag, None, None, None)) + continue + med, arr, peak = out + lp = _lpips(ref, arr) + rows.append((tag, med, peak, lp)) + spd = f"{bmed/med:.2f}x" if med else "-" + print(f" {tag:20s} {med:.3f}s ({spd} vs base) peak={peak:.1f}G LPIPS={lp}", flush = True) + + print("\n==== SUMMARY (ref = baseline compile) ====", flush = True) + for tag, med, peak, lp in rows: + if med is None: + print(f" {tag:20s} FAILED") + continue + spd = f"{bmed/med:.2f}x" if med else "-" + lpv = "ref" if (tag == "baseline") else (f"{lp:.3f}" if lp is not None else "n/a") + print(f" {tag:20s} {med:.3f}s {spd:>6s} peak={peak:.1f}G LPIPS={lpv:>6s}", flush = True) + print("PERF-LEVERS-DONE", flush = True) + return 0 + + +if __name__ == "__main__": + sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "studio" / "backend")) + sys.exit(main()) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index aac4069833..b61badb50b 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -52,6 +52,10 @@ from .diffusion_speed import ( restore_backend_flags, snapshot_backend_flags, ) +from .diffusion_attention import ( + apply_attention_backend, + select_attention_backend, +) from .diffusion_precision import quantize_text_encoders from .diffusion_prequant import ( load_prequantized_transformer, @@ -96,6 +100,9 @@ class _LoadState: # 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 (e.g. "_native_cudnn"), or + # None for the default SDPA. Set before compile; orthogonal to the weight quant. + attention_backend: Optional[str] = None @dataclass @@ -306,6 +313,7 @@ class DiffusionBackend: transformer_quant: Optional[str] = None, transformer_quant_fast_accum: Optional[bool] = None, transformer_prequant_path: Optional[str] = None, + attention_backend: Optional[str] = None, ) -> dict[str, Any]: """Validate, then run the (slow) load on a daemon thread. Returns at once.""" fam = self.validate_load_request( @@ -340,6 +348,7 @@ class DiffusionBackend: transformer_quant = transformer_quant, transformer_quant_fast_accum = transformer_quant_fast_accum, transformer_prequant_path = transformer_prequant_path, + attention_backend = attention_backend, _load_token = token, ), daemon = True, @@ -476,6 +485,7 @@ class DiffusionBackend: transformer_quant: Optional[str] = None, transformer_quant_fast_accum: Optional[bool] = None, transformer_prequant_path: Optional[str] = None, + attention_backend: Optional[str] = None, _load_token: Optional[int] = None, ) -> dict[str, Any]: # Validate first (cheap, no torch/diffusers) so a direct call with a bad @@ -596,6 +606,18 @@ class DiffusionBackend: # first so unload can restore them: TF32 / cudnn.benchmark are global, # and a later `off` load must not inherit this load's settings. backend_flags_before = snapshot_backend_flags() + # Pick the attention kernel BEFORE compile (compile traces attention). auto + # upgrades to cuDNN fused attention on NVIDIA when a speed profile is active + # (~1.18x, near-lossless); an explicit backend is honored, falling back to + # the diffusers default if its kernel is unavailable. Orthogonal to the + # weight quant -- it speeds the QK/PV matmuls torchao does not touch. + attention_engaged = apply_attention_backend( + pipe, + select_attention_backend( + target, attention_backend, speed_active = effective_speed != SPEED_OFF + ), + logger = logger, + ) speed_applied = apply_speed_optims( pipe, target, @@ -649,6 +671,7 @@ class DiffusionBackend: backend_flags_before = backend_flags_before, text_encoder_quant = te_quant, transformer_quant = transformer_quant_engaged, + attention_backend = attention_engaged, ) logger.info( @@ -953,6 +976,7 @@ class DiffusionBackend: "speed_optims": [], "text_encoder_quant": None, "transformer_quant": None, + "attention_backend": None, } return { "loaded": True, @@ -969,6 +993,7 @@ class DiffusionBackend: "speed_optims": list(state.speed_optims), "text_encoder_quant": state.text_encoder_quant, "transformer_quant": state.transformer_quant, + "attention_backend": state.attention_backend, } diff --git a/studio/backend/core/inference/diffusion_attention.py b/studio/backend/core/inference/diffusion_attention.py new file mode 100644 index 0000000000..e652b0068c --- /dev/null +++ b/studio/backend/core/inference/diffusion_attention.py @@ -0,0 +1,245 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Select the diffusion transformer's attention backend. + +diffusers exposes a unified ``transformer.set_attention_backend(name)`` dispatcher that +swaps the scaled-dot-product-attention kernel, validating hardware/package requirements at +set time and otherwise leaving the default (``native`` = ``F.scaled_dot_product_attention``). +Attention is memory-bandwidth bound, so a better kernel is a real end-to-end win that is +orthogonal to the linear-weight quantisation (it speeds the QK/PV matmuls torchao never +touches) and composes with torch.compile. + + auto - the best *exact* (non-quantized) backend for the device. On NVIDIA CUDA that is + cuDNN's fused attention (``_native_cudnn``), measured ~1.18x end-to-end on a B200 + with LPIPS ~0.004 vs the default (below the compile/quant noise floor). On + AMD/Intel/Apple/CPU it stays ``native`` (the dispatcher already routes those). + ``auto`` only upgrades when a speed profile is active, so ``speed_mode=off`` stays + bit-identical. + native - force the default SDPA (bit-identical reference). + cudnn - cuDNN fused attention (exact; NVIDIA). + flash / flash3 / flash4 - FlashAttention 2 / 3 (Hopper) / 4 (SM100); exact, kernel-gated. + sage - SageAttention (INT8 QK); quantized, a small quality cost, consumer-friendly. + xformers / aiter - memory-efficient (NVIDIA) / AITER (AMD ROCm). + +Best-effort: an unavailable backend (missing kernel / wrong arch) is caught and the load +falls back to the diffusers default rather than failing. torch/diffusers imported lazily. +""" + +from __future__ import annotations + +from typing import Any, Optional + +ATTN_AUTO = "auto" +ATTN_NATIVE = "native" + +# User-facing alias -> the diffusers dispatcher backend name. +_ALIASES: dict[str, str] = { + "native": "native", + "sdpa": "native", + "cudnn": "_native_cudnn", + "flash": "flash", + "flash2": "flash", + "flash3": "_flash_3_hub", + "flash4": "flash_4_hub", + "sage": "sage", + "xformers": "xformers", + "aiter": "aiter", +} +ATTN_ALIASES = (ATTN_AUTO,) + tuple(dict.fromkeys(_ALIASES)) + + +def normalize_attention_backend(value: Optional[str]) -> Optional[str]: + """Lower/strip a requested attention backend; None / "" / "auto" -> "auto". + + Raises ValueError for an unsupported alias so a bad request is rejected cheaply.""" + if value is None: + return ATTN_AUTO + normalized = str(value).strip().lower() + if not normalized: + return ATTN_AUTO + if normalized not in ATTN_ALIASES: + raise ValueError( + f"Unsupported attention_backend '{value}'. Use one of: {', '.join(ATTN_ALIASES)}." + ) + return normalized + + +# Backends diffusers validates only by *package* at set time (``_check_attention_backend_ +# requirements`` checks the ``kernels`` install, not the GPU), but whose kernels need a +# specific CUDA arch at run time -- so an explicit request on the wrong card loads/sets fine +# and then crashes mid-generation. Gate them up front by a (min, max-exclusive) compute +# capability range. FlashAttention 3 is a Hopper-SM90 rewrite with no Blackwell kernel, so it +# needs an upper bound: an explicit flash3 on a B200 (SM100) must drop to native instead of +# setting fine then crashing at generation. FlashAttention 4 is Blackwell+ (no upper bound). +_ARCH_CAPABILITY: dict[str, tuple[tuple[int, int], Optional[tuple[int, int]]]] = { + "_flash_3_hub": ((9, 0), (10, 0)), # FlashAttention 3 -> Hopper (SM90) only + "flash_4_hub": ((10, 0), None), # FlashAttention 4 -> Blackwell (SM100)+ +} + + +def _cuda_capability() -> Optional[tuple[int, int]]: + """(major, minor) compute capability of the active CUDA device, or None if unknown.""" + try: + import torch + if not torch.cuda.is_available(): + return None + return tuple(torch.cuda.get_device_capability()) # type: ignore[return-value] + except Exception: # noqa: BLE001 + return None + + +def _backend_arch_supported(backend: str) -> bool: + """False only when ``backend`` needs a CUDA arch outside this device's supported range. + + Unknown capability (no CUDA / detection failure) returns True so we never block on a + guess -- diffusers' own set-time check still guards the package, and a genuine run-time + failure falls back to native.""" + bounds = _ARCH_CAPABILITY.get(backend) + if bounds is None: + return True + have = _cuda_capability() + if have is None: + return True + low, high = bounds + return have >= low and (high is None or have < high) + + +def _is_cuda_nvidia(target: Any) -> bool: + """CUDA device on an NVIDIA (non-ROCm) build -- where cuDNN attention applies.""" + if getattr(target, "device", None) != "cuda": + return False + try: + import torch + return getattr(torch.version, "hip", None) is None + except Exception: # noqa: BLE001 + return False + + +def select_attention_backend( + target: Any, requested: Optional[str], *, speed_active: bool +) -> Optional[str]: + """The dispatcher backend name to apply, or None to leave the diffusers default. + + An explicit alias is honored verbatim (apply falls back if its kernel is unavailable). + ``auto`` upgrades to cuDNN on NVIDIA CUDA only when a speed profile is active (so + ``off`` stays bit-identical); everywhere else it returns None (native default).""" + alias = normalize_attention_backend(requested) + if alias != ATTN_AUTO: + backend = _ALIASES[alias] + if backend == "native": + return None + # An arch-gated kernel (flash3/flash4) on a card that can't run it would set fine + # then crash mid-generation, so drop it to the native default up front. + if not _backend_arch_supported(backend): + return None + # cuDNN fused SDPA needs Ampere+ (SM80); diffusers accepts it on pre-SM80 cards + # (T4/V100) then fails at the first generation, so apply the same gate to an + # explicit cuDNN request as the auto path already does. + if backend == "_native_cudnn" and not _cudnn_attention_supported(): + return None + return backend + # auto + if speed_active and _is_cuda_nvidia(target) and _cudnn_attention_supported(): + return "_native_cudnn" + return None + + +def _cudnn_attention_supported() -> bool: + """cuDNN fused SDPA needs Ampere+ (SM80). On pre-SM80 NVIDIA cards (T4 SM75 / + V100 SM70) diffusers accepts ``_native_cudnn`` at set time but the kernel fails at + the first generation, so gate the auto-cuDNN upgrade on capability. Unknown + capability allows it (diffusers' set-time check + the run-time fallback still guard).""" + have = _cuda_capability() + return have is None or have >= (8, 0) + + +def apply_attention_backend( + pipe: Any, + backend: Optional[str], + *, + logger: Any = None, +) -> Optional[str]: + """Set ``backend`` on ``pipe.transformer`` via the diffusers dispatcher. + + Returns the backend actually engaged, or None when left at the native default (either + because ``backend`` was None or because the requested kernel was unavailable -> graceful + fallback, never a load failure). + + diffusers keeps a *process-wide* active attention backend that ``set_attention_backend`` + also updates, and a fresh transformer's processors follow it (their ``_attention_backend`` + defaults to None). So a load that wants native must restore it explicitly: otherwise it + silently inherits a backend an earlier load pinned (e.g. cuDNN under a speed profile), + breaking the bit-identical/``off`` guarantee. Best-effort throughout.""" + transformer = getattr(pipe, "transformer", None) + fn = getattr(transformer, "set_attention_backend", None) + if not callable(fn): + return None + if backend is not None: + try: + fn(backend) + # set_attention_backend also pins the backend in diffusers' process-wide + # registry. This transformer's own processors keep it locally (their + # _attention_backend is now explicit), so reset the global default back to + # native -- otherwise a later component whose processors are unconfigured + # (backend None) silently inherits this kernel. + _reset_global_backend_to_native(logger) + if logger is not None: + logger.info("diffusion.attention: backend=%s", backend) + return backend + except Exception as exc: # noqa: BLE001 — unavailable kernel -> restore native below + _warn(logger, backend, exc) + # No backend requested, or the requested one failed: pin the native default so a stale + # process-wide backend from a previous load can't leak into this one. + _restore_native_backend(fn, logger) + return None + + +def _active_attention_backend() -> Optional[str]: + """The diffusers process-wide active attention backend name, or None if undeterminable.""" + try: + from diffusers.models.attention_dispatch import _AttentionBackendRegistry + + # get_active_backend() returns a (AttentionBackendName, fn) tuple (or None), so + # take element 0 and read its .value (e.g. "native"); reading .value off the + # tuple itself would yield a junk string that never compares equal to a name. + active = _AttentionBackendRegistry.get_active_backend() + if active is None: + return None + name = active[0] if isinstance(active, tuple) else active + return getattr(name, "value", str(name)) + except Exception: # noqa: BLE001 + return None + + +def _reset_global_backend_to_native(logger: Any) -> None: + """Reset diffusers' process-wide active attention backend to native after a + successful per-transformer set, so a later component whose processors are + unconfigured (backend None) does not inherit this transformer's kernel. The + transformer's own processors keep the backend just set. Best-effort and silent: + if the diffusers internals move, the prior (leaking) behavior is unchanged.""" + if _active_attention_backend() == ATTN_NATIVE: + return + try: + from diffusers.models.attention_dispatch import ( + AttentionBackendName, + _AttentionBackendRegistry, + ) + _AttentionBackendRegistry.set_active_backend(AttentionBackendName.NATIVE) + except Exception: # noqa: BLE001 — best-effort; leave the global as-is on any change + pass + + +def _restore_native_backend(set_backend_fn: Any, logger: Any) -> None: + """Force the native default when the global active backend isn't already native.""" + if _active_attention_backend() == ATTN_NATIVE: + return # already native -> avoid redundant work and an extra dispatcher warning + try: + set_backend_fn(ATTN_NATIVE) + except Exception as exc: # noqa: BLE001 — best-effort restore + _warn(logger, ATTN_NATIVE, exc) + + +def _warn(logger: Any, what: str, exc: Exception) -> None: + if logger is not None: + logger.warning("diffusion.attention: %s unavailable (%s); using default", what, exc) diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 08d309f257..402c931da1 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -1750,6 +1750,30 @@ class DiffusionLoadRequest(BaseModel): "the OS path separator). A bare on/off value such as '1' is deliberately not " "accepted -- it must name an allowed directory.", ) + attention_backend: Optional[ + Literal[ + "auto", + "native", + "sdpa", + "cudnn", + "flash", + "flash2", + "flash3", + "flash4", + "sage", + "xformers", + "aiter", + ] + ] = Field( + None, + description = "Attention kernel via the diffusers dispatcher. auto picks the best " + "exact backend for the device (cuDNN fused attention on NVIDIA, ~1.18x and " + "near-lossless, when a speed profile is active; native SDPA elsewhere and when " + "speed=off). native (alias sdpa) forces default SDPA; cudnn/flash/flash3/flash4 are exact " + "(kernel/arch-gated); sage is INT8 attention (a small quality cost, consumer " + "friendly); xformers/aiter are memory-efficient (NVIDIA) / AMD ROCm. An " + "unavailable kernel falls back to the default.", + ) class DiffusionGenerateRequest(BaseModel): @@ -1865,3 +1889,8 @@ class DiffusionStatusResponse(BaseModel): description = "Transformer quant engaged on the dense fast path: int8 | fp8 | " "nvfp4 | mxfp8 | null (null = the GGUF transformer was loaded)", ) + attention_backend: Optional[str] = Field( + None, + description = "Attention backend engaged via the diffusers dispatcher (e.g. " + "_native_cudnn), or null for the default SDPA", + ) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index ec5382e190..c89a0cc7eb 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -10334,6 +10334,7 @@ async def load_diffusion_model( transformer_quant = request.transformer_quant, transformer_quant_fast_accum = request.transformer_quant_fast_accum, transformer_prequant_path = request.transformer_prequant_path, + attention_backend = request.attention_backend, ) return DiffusionStatusResponse(**status_dict) except (ValueError, FileNotFoundError) as exc: diff --git a/studio/backend/tests/test_diffusion_attention.py b/studio/backend/tests/test_diffusion_attention.py new file mode 100644 index 0000000000..3e43aab8a9 --- /dev/null +++ b/studio/backend/tests/test_diffusion_attention.py @@ -0,0 +1,225 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Hermetic CPU tests for attention-backend selection. No torch/diffusers needed: +``_is_cuda_nvidia`` is monkeypatched for the policy tests, and the apply path uses a fake +transformer that records / raises on ``set_attention_backend``. +""" + +from __future__ import annotations + +import types + +import pytest + +import core.inference.diffusion_attention as att +from core.inference.diffusion_attention import ( + ATTN_AUTO, + apply_attention_backend, + normalize_attention_backend, + select_attention_backend, +) + + +def _target(device = "cuda"): + return types.SimpleNamespace(device = device) + + +# ── normalize ──────────────────────────────────────────────────────────────────── +def test_normalize_defaults_and_aliases(): + assert normalize_attention_backend(None) == ATTN_AUTO + assert normalize_attention_backend("") == ATTN_AUTO + assert normalize_attention_backend("auto") == ATTN_AUTO + assert normalize_attention_backend("CuDNN") == "cudnn" + assert normalize_attention_backend("FLASH3") == "flash3" + assert normalize_attention_backend("sdpa") == "sdpa" + + +def test_normalize_rejects_unknown(): + with pytest.raises(ValueError): + normalize_attention_backend("bogus") + # dashes are no longer silently rewritten to underscores -> a dashed alias is rejected. + with pytest.raises(ValueError): + normalize_attention_backend("flash-3") + + +def test_sdpa_alias_maps_to_native(): + # sdpa is an alias for native -> nothing to set on the dispatcher. + assert select_attention_backend(_target(), "sdpa", speed_active = True) is None + + +# ── select policy ───────────────────────────────────────────────────────────────── +def test_auto_upgrades_to_cudnn_on_nvidia_when_speed_active(monkeypatch): + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True) + monkeypatch.setattr(att, "_cuda_capability", lambda: (8, 0)) # Ampere+: cuDNN ok + assert select_attention_backend(_target(), "auto", speed_active = True) == "_native_cudnn" + + +def test_auto_does_not_pin_cudnn_below_sm80(monkeypatch): + # cuDNN fused SDPA fails at run time on pre-SM80 (T4 SM75 / V100 SM70); auto must stay + # on the native default there rather than pin a backend that crashes on first generation. + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True) + monkeypatch.setattr(att, "_cuda_capability", lambda: (7, 5)) # Turing T4 + assert select_attention_backend(_target(), "auto", speed_active = True) is None + + +def test_auto_stays_native_when_speed_off(monkeypatch): + # off must stay bit-identical -> no backend change even on NVIDIA. + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True) + assert select_attention_backend(_target(), "auto", speed_active = False) is None + + +def test_auto_stays_native_off_nvidia(monkeypatch): + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False) + assert select_attention_backend(_target(device = "mps"), "auto", speed_active = True) is None + + +def test_explicit_backend_honored_regardless_of_speed(monkeypatch): + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False) + # Pin a high capability so the arch-gated flash4 isn't dropped by the runtime check. + monkeypatch.setattr(att, "_cuda_capability", lambda: (10, 0)) + assert select_attention_backend(_target(), "sage", speed_active = False) == "sage" + assert select_attention_backend(_target(), "flash4", speed_active = False) == "flash_4_hub" + assert select_attention_backend(_target(), "cudnn", speed_active = False) == "_native_cudnn" + + +def test_explicit_native_returns_none(): + # native is the default -> nothing to set. + assert select_attention_backend(_target(), "native", speed_active = True) is None + + +# ── arch gating (flash3/flash4 need a specific CUDA capability) ───────────────────── +def test_flash3_dropped_below_hopper(monkeypatch): + monkeypatch.setattr(att, "_cuda_capability", lambda: (8, 9)) # Ada / consumer + assert select_attention_backend(_target(), "flash3", speed_active = False) is None + + +def test_flash4_dropped_below_blackwell(monkeypatch): + monkeypatch.setattr(att, "_cuda_capability", lambda: (9, 0)) # Hopper, but FA4 needs SM100 + assert select_attention_backend(_target(), "flash4", speed_active = False) is None + # flash3 still allowed on Hopper. + assert select_attention_backend(_target(), "flash3", speed_active = False) == "_flash_3_hub" + + +def test_arch_gate_does_not_block_when_capability_unknown(monkeypatch): + # Unknown capability (e.g. no CUDA) must not block -> diffusers' set-time check still guards. + monkeypatch.setattr(att, "_cuda_capability", lambda: None) + assert select_attention_backend(_target(), "flash4", speed_active = False) == "flash_4_hub" + + +def test_flash3_dropped_on_blackwell(monkeypatch): + # FlashAttention 3 is a Hopper-SM90 rewrite with no Blackwell kernel: an explicit + # flash3 on a B200 (SM100) must drop to native rather than set fine then crash. + monkeypatch.setattr(att, "_cuda_capability", lambda: (10, 0)) + assert select_attention_backend(_target(), "flash3", speed_active = False) is None + # FA4 is still honored on Blackwell. + assert select_attention_backend(_target(), "flash4", speed_active = False) == "flash_4_hub" + # flash3 is allowed exactly on Hopper SM90. + monkeypatch.setattr(att, "_cuda_capability", lambda: (9, 0)) + assert select_attention_backend(_target(), "flash3", speed_active = False) == "_flash_3_hub" + + +def test_explicit_cudnn_dropped_below_sm80(monkeypatch): + # An explicit cuDNN request on pre-Ampere (T4 SM75 / V100 SM70) must drop to native, + # not set fine and crash at first generation -- the same gate the auto path applies. + monkeypatch.setattr(att, "_cuda_capability", lambda: (7, 5)) + assert select_attention_backend(_target(), "cudnn", speed_active = False) is None + # Ampere+ still honors it. + monkeypatch.setattr(att, "_cuda_capability", lambda: (8, 0)) + assert select_attention_backend(_target(), "cudnn", speed_active = False) == "_native_cudnn" + + +# ── apply ───────────────────────────────────────────────────────────────────────── +class _FakeTransformer: + def __init__(self, *, fail = False): + self.fail = fail + self.set_to = None + + def set_attention_backend(self, name): + if self.fail: + raise RuntimeError(f"{name} kernel unavailable") + self.set_to = name + + +def _pipe(transformer): + return types.SimpleNamespace(transformer = transformer) + + +def test_apply_none_leaves_native_when_global_already_native(monkeypatch): + # Global already native -> no redundant set call, returns None. + monkeypatch.setattr(att, "_active_attention_backend", lambda: "native") + t = _FakeTransformer() + assert apply_attention_backend(_pipe(t), None) is None + assert t.set_to is None + + +def test_apply_none_restores_native_when_global_polluted(monkeypatch): + # A previous load pinned cuDNN process-wide; a native load must reset it so it can't + # silently inherit cuDNN (the bit-identical/off guarantee). + monkeypatch.setattr(att, "_active_attention_backend", lambda: "_native_cudnn") + t = _FakeTransformer() + assert apply_attention_backend(_pipe(t), None) is None + assert t.set_to == "native" + + +def test_apply_sets_backend(): + t = _FakeTransformer() + engaged = apply_attention_backend(_pipe(t), "_native_cudnn") + assert engaged == "_native_cudnn" and t.set_to == "_native_cudnn" + + +def test_apply_falls_back_on_unavailable_kernel(monkeypatch): + # an unavailable kernel must not fail the load -> returns None (diffusers default). + monkeypatch.setattr(att, "_active_attention_backend", lambda: "native") + t = _FakeTransformer(fail = True) + assert apply_attention_backend(_pipe(t), "sage") is None + + +def test_apply_failed_kernel_restores_native_when_polluted(monkeypatch): + # Requested kernel fails AND the global is polluted: restore native before returning. + monkeypatch.setattr(att, "_active_attention_backend", lambda: "_native_cudnn") + + class _FailOnceTransformer: + def __init__(self): + self.calls = [] + + def set_attention_backend(self, name): + self.calls.append(name) + if name != "native": + raise RuntimeError(f"{name} kernel unavailable") + + t = _FailOnceTransformer() + assert apply_attention_backend(_pipe(t), "sage") is None + assert t.calls == ["sage", "native"] + + +def test_apply_handles_missing_method(): + pipe = types.SimpleNamespace(transformer = types.SimpleNamespace()) + assert apply_attention_backend(pipe, "_native_cudnn") is None + + +def test_apply_resets_global_registry_after_success(monkeypatch): + # After a successful per-transformer set, the process-wide registry must be reset to + # native so a later component (unconfigured processors) can't inherit this kernel -- + # while the transformer's own backend stays the engaged one. + called = {"reset": False} + monkeypatch.setattr( + att, "_reset_global_backend_to_native", lambda logger: called.__setitem__("reset", True) + ) + t = _FakeTransformer() + engaged = apply_attention_backend(_pipe(t), "_native_cudnn") + assert engaged == "_native_cudnn" and t.set_to == "_native_cudnn" + assert called["reset"] is True + + +def test_active_attention_backend_reads_tuple_return(): + # get_active_backend() returns a (AttentionBackendName, fn) tuple; the helper must read + # the name's .value, not stringify the tuple (which never compares equal to a name). + pytest.importorskip("diffusers") + from diffusers.models.attention_dispatch import ( + AttentionBackendName, + _AttentionBackendRegistry, + ) + + _AttentionBackendRegistry.set_active_backend(AttentionBackendName.NATIVE) + assert att._active_attention_backend() == "native" diff --git a/studio/backend/tests/test_diffusion_routes.py b/studio/backend/tests/test_diffusion_routes.py index 839d325443..dc34e60cef 100644 --- a/studio/backend/tests/test_diffusion_routes.py +++ b/studio/backend/tests/test_diffusion_routes.py @@ -424,6 +424,29 @@ def test_transformer_prequant_path_threads_through(client, monkeypatch): assert backend.last_load_kwargs.get("transformer_prequant_path") == "/data/zimage_fp8.pt" +def test_attention_backend_threads_through(client, monkeypatch): + backend = _FakeBackend() + monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend) + resp = client.post( + "/api/inference/images/load", + json = { + "model_path": "x/z-image", + "gguf_filename": "q.gguf", + "attention_backend": "cudnn", + }, + ) + assert resp.status_code == 200 + assert backend.last_load_kwargs.get("attention_backend") == "cudnn" + + +def test_invalid_attention_backend_returns_422(client): + resp = client.post( + "/api/inference/images/load", + json = {"model_path": "x/z-image", "gguf_filename": "q.gguf", "attention_backend": "bogus"}, + ) + assert resp.status_code == 422 + + def test_prequant_path_doc_describes_allowlist_not_toggle(): # The field help must match the code: UNSLOTH_ALLOW_LOCAL_PREQUANT_PATH is a # directory allowlist, not a =1 toggle (diffusion_prequant._allowed_prequant_roots