Studio diffusion (Phase 10): attention-backend selection

Add a selectable attention kernel via the diffusers set_attention_backend
dispatcher. Attention is memory-bandwidth bound, so a better kernel is an
end-to-end win orthogonal to the linear-weight quantisation (it speeds the QK/PV
matmuls torchao never touches) and composes with torch.compile.

auto picks the best exact backend for the device: cuDNN fused attention
(_native_cudnn) on NVIDIA when a speed profile is active, measured ~1.18x
end-to-end on a B200 (Z-Image 1024px/8 steps) with LPIPS ~0.004 vs the default
(below the compile/quant noise floor); native SDPA elsewhere and when speed=off
(so off stays bit-identical). Explicit native/cudnn/flash/flash3/flash4/sage/
xformers/aiter are honored, and an unavailable kernel falls back to the default
rather than failing the load.

New core/inference/diffusion_attention.py (normalize + per-device select + apply,
best-effort, lazy imports). Set on pipe.transformer BEFORE compile in load_pipeline;
attention_backend threads through begin_load / load_pipeline / status like the other
load knobs. New request field attention_backend + status field. Hermetic CPU tests
for normalize / select policy / apply fallback, plus route threading + 422. Measured
via scripts/perf_levers_probe.py.
This commit is contained in:
Daniel Han 2026-06-26 12:02:27 +00:00
commit b923675549
7 changed files with 490 additions and 0 deletions

View file

@ -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())

View file

@ -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,
}

View file

@ -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)

View file

@ -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",
)

View file

@ -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:

View file

@ -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

View file

@ -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.