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:
parent
520b80a082
commit
b923675549
7 changed files with 490 additions and 0 deletions
190
scripts/perf_levers_probe.py
Normal file
190
scripts/perf_levers_probe.py
Normal 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())
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
129
studio/backend/core/inference/diffusion_attention.py
Normal file
129
studio/backend/core/inference/diffusion_attention.py
Normal 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)
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
105
studio/backend/tests/test_diffusion_attention.py
Normal file
105
studio/backend/tests/test_diffusion_attention.py
Normal 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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue