diff --git a/scripts/perf_levers_probe.py b/scripts/perf_levers_probe.py new file mode 100644 index 0000000000..97b0daafa7 --- /dev/null +++ b/scripts/perf_levers_probe.py @@ -0,0 +1,190 @@ +# 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("/mnt/disks/unslothai/ubuntu/workspace_81/outputs/quant_research/perf_levers_images") + + +_LP = {"fn": None} + + +def _lpips(ref, arr): + try: + import lpips + import torch + if _LP["fn"] is None: + _LP["fn"] = lpips.LPIPS(net="alex", verbose=False).cuda().eval() + + def t(x): + return (torch.from_numpy(x).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0).cuda() + + 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 + + +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) + 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) + 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 cbb5abd2ab..126071c10a 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -51,6 +51,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, @@ -94,6 +98,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 @@ -283,6 +290,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( @@ -317,6 +325,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, @@ -441,6 +450,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 @@ -549,6 +559,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, @@ -592,6 +614,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( @@ -893,6 +916,7 @@ class DiffusionBackend: "speed_optims": [], "text_encoder_quant": None, "transformer_quant": None, + "attention_backend": None, } return { "loaded": True, @@ -909,6 +933,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..8725046791 --- /dev/null +++ b/studio/backend/core/inference/diffusion_attention.py @@ -0,0 +1,129 @@ +# 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().replace("-", "_") + 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 + + +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] + return None if backend == "native" else backend + # auto + if speed_active and _is_cuda_nvidia(target): + return "_native_cudnn" + return None + + +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 default (either because + ``backend`` was None or because the requested kernel was unavailable -> graceful + fallback to the diffusers default, never a load failure). Best-effort.""" + if backend is None: + return None + transformer = getattr(pipe, "transformer", None) + fn = getattr(transformer, "set_attention_backend", None) + if not callable(fn): + return None + try: + fn(backend) + if logger is not None: + logger.info("diffusion.attention: backend=%s", backend) + return backend + except Exception as exc: # noqa: BLE001 — unavailable kernel -> diffusers default + _warn(logger, backend, exc) + return None + + +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 928df30358..005376086f 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -1746,6 +1746,18 @@ class DiffusionLoadRequest(BaseModel): "GPU (~half the load VRAM and a smaller download). null uses the family's hosted " "checkpoint if configured, else quantises the dense transformer at load time.", ) + attention_backend: Optional[ + Literal["auto", "native", "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 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): @@ -1858,3 +1870,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 06df9eb154..4dca8e262d 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -10088,6 +10088,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..27dbf1e32c --- /dev/null +++ b/studio/backend/tests/test_diffusion_attention.py @@ -0,0 +1,105 @@ +# 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" + + +def test_normalize_rejects_unknown(): + with pytest.raises(ValueError): + normalize_attention_backend("bogus") + + +# ── select policy ───────────────────────────────────────────────────────────────── +def test_auto_upgrades_to_cudnn_on_nvidia_when_speed_active(monkeypatch): + monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True) + assert select_attention_backend(_target(), "auto", speed_active=True) == "_native_cudnn" + + +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) + 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 + + +# ── 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_is_noop(): + assert apply_attention_backend(_pipe(_FakeTransformer()), None) is None + + +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(): + # an unavailable kernel must not fail the load -> returns None (diffusers default). + t = _FakeTransformer(fail=True) + assert apply_attention_backend(_pipe(t), "sage") is None + + +def test_apply_handles_missing_method(): + pipe = types.SimpleNamespace(transformer=types.SimpleNamespace()) + assert apply_attention_backend(pipe, "_native_cudnn") is None diff --git a/studio/backend/tests/test_diffusion_routes.py b/studio/backend/tests/test_diffusion_routes.py index cd41397901..cf1d4266b6 100644 --- a/studio/backend/tests/test_diffusion_routes.py +++ b/studio/backend/tests/test_diffusion_routes.py @@ -359,6 +359,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_invalid_transformer_quant_returns_422_without_eviction(client): # An unsupported transformer_quant is rejected by the request schema (Literal), so # the GPU is never acquired and no chat model is evicted.