unsloth/studio/backend/core/inference/diffusion_attention.py
Daniel Han 36df317293 Trim the comments across the diffusion backend
Comment-only pass over the Python this PR touches: drop what the code already
says, collapse multi-line explanations that still read on one line, and keep
the reasoning that is not recoverable from the code. No code, docstring
semantics or behaviour changes; verified with an AST comparison against the
previous revision, and the backend suite is unchanged (same 37 environment
failures as before: the API integration tests that need a live keyed server,
the flash-attn install hooks, and the GPU memory fields).
2026-07-26 20:31:19 +00:00

344 lines
14 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Select the diffusion transformer's attention backend.
diffusers' ``transformer.set_attention_backend(name)`` dispatcher swaps the SDPA kernel,
validating hardware/package at set time (default ``native`` = ``F.scaled_dot_product_attention``).
Attention is bandwidth-bound, so a better kernel is a real win orthogonal to weight quantisation
(it speeds the QK/PV matmuls torchao never touches) and composes with torch.compile.
auto - the best *exact* backend for the device. On NVIDIA CUDA that is cuDNN fused attention
(``_native_cudnn``), ~1.18x end-to-end on B200, LPIPS ~0.004 (below the noise floor).
Elsewhere stays ``native``. 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, small quality cost, consumer-friendly.
xformers / aiter - memory-efficient (NVIDIA) / AITER (AMD ROCm).
Best-effort: an unavailable backend falls back to the diffusers default. torch/diffusers lazy.
"""
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 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 but whose kernels need a specific CUDA
# arch at run time. Gate by a (min, max-exclusive) capability range: FA3 is Hopper-SM90 only (so
# flash3 on a B200 drops to native), FA4 is Blackwell+.
_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 range. Unknown
capability returns True (never block on a guess; the 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 (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); elsewhere returns None (native)."""
alias = normalize_attention_backend(requested)
if alias != ATTN_AUTO:
backend = _ALIASES[alias]
if backend == "native":
return None
# AITER is the AMD ROCm kernel: honor it on a ROCm target, drop it elsewhere (else the
# NVIDIA-only guard below would drop the one backend that only works on ROCm).
if backend == "aiter":
if getattr(target, "device", None) == "cuda" and not _is_cuda_nvidia(target):
return backend
return None
# cuDNN / flash* / sage are CUDA+NVIDIA-only; elsewhere the first generation crashes.
if not _is_cuda_nvidia(target):
return None
# An arch-gated kernel (flash3/flash4) on a card that can't run it sets fine then crashes.
if not _backend_arch_supported(backend):
return None
# cuDNN fused SDPA needs Ampere+ (SM80); gate an explicit request like the auto path.
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 cards (T4/V100) diffusers accepts it
then fails at generation, so gate the upgrade on capability. Unknown capability allows it."""
have = _cuda_capability()
return have is None or have >= (8, 0)
# Optional-kernel backends installable on demand: dispatcher name -> (probe module, pip package).
# Wheels only (--only-binary=:all:): a source build needs a CUDA toolchain a Studio host may
# lack. cuDNN/native ship with torch.
_INSTALLABLE_BACKENDS: dict[str, tuple[str, str]] = {
"sage": ("sageattention", "sageattention"),
"flash": ("flash_attn", "flash-attn"),
"_flash_3_hub": ("kernels", "kernels"), # FA3/FA4 from the HF kernels hub
"flash_4_hub": ("kernels", "kernels"),
"xformers": ("xformers", "xformers"),
}
# On-demand install gate (mirrors UNSLOTH_DIFFUSION_SD_CPP_INSTALL):
# auto (default) / 1 - install the missing package when a gated backend is requested
# 0 - never install; a missing kernel falls back to native
_ATTENTION_INSTALL_ENV = "UNSLOTH_DIFFUSION_ATTENTION_INSTALL"
# Packages a pip install was already attempted for in THIS process. The loader pre-installs
# outside its locks, then re-resolves under _generate_lock where apply would otherwise re-run
# the full 600s install holding the load lock; a recorded attempt makes the retry a no-op.
_INSTALL_ATTEMPTED: set[str] = set()
def _ensure_attention_backend_installed(backend: str, logger: Any = None) -> None:
"""Best-effort wheel-only install of the package ``backend`` needs, when allowed.
Called after arch gating, so only for a backend that could work here. Failure is swallowed:
the subsequent set_attention_backend raises on the missing package and falls back to native."""
import importlib.util
import os
spec = _INSTALLABLE_BACKENDS.get(backend)
if spec is None:
return
module, package = spec
gate = os.environ.get(_ATTENTION_INSTALL_ENV, "auto").strip().lower()
if gate in ("0", "false", "no", "off"):
return
try:
if importlib.util.find_spec(module) is not None:
return
except Exception: # noqa: BLE001 — a broken install probes as missing; try the install
pass
# Attempt each install once per process, else the in-lock apply path re-runs the whole install
# under _generate_lock and blocks unload/cancel.
if package in _INSTALL_ATTEMPTED:
return
_INSTALL_ATTEMPTED.add(package)
import subprocess
import sys
if logger is not None:
logger.info(
"diffusion.attention: installing %s for backend=%s (wheel-only)", package, backend
)
try:
subprocess.run(
# --no-deps: install ONLY this kernel wheel, since xformers/flash-attn pin an exact torch and
# normal resolution would replace the running one. An ABI mismatch just fails to import.
[
sys.executable,
"-m",
"pip",
"install",
"--only-binary",
":all:",
"--no-deps",
package,
],
capture_output = True,
timeout = 600,
check = True,
)
# The import system caches directory listings, so invalidate the finder caches or the next
# find_spec can miss the wheel just installed.
importlib.invalidate_caches()
except Exception as exc: # noqa: BLE001 — no wheel / no network -> native fallback
if logger is not None:
# CalledProcessError.str() shows only the exit code; surface stderr so the fallback is
# diagnosable.
stderr = getattr(exc, "stderr", None)
if stderr:
if isinstance(stderr, bytes):
stderr = stderr.decode("utf-8", errors = "replace")
logger.warning(
"diffusion.attention: could not install %s; pip failed with: %s",
package,
stderr.strip() or str(exc),
)
else:
logger.warning(
"diffusion.attention: could not install %s (%s); falling back to default",
package,
exc,
)
def _attention_dits(pipe: Any) -> list:
"""Every DiT the denoise loop runs: the primary ``transformer`` plus a second expert some
families carry (Ideogram's ``unconditional_transformer``, an MoE ``transformer_2``). The
backend must be set on ALL of them, else the second DiT keeps the native default."""
dits: list = []
for attr in ("transformer", "transformer_2", "unconditional_transformer"):
m = getattr(pipe, attr, None)
if m is not None and m not in dits:
dits.append(m)
return dits
def apply_attention_backend(
pipe: Any,
backend: Optional[str],
*,
logger: Any = None,
) -> Optional[str]:
"""Set ``backend`` on EVERY denoiser DiT via the diffusers dispatcher.
Returns the backend engaged, or None when left at native (``backend`` was None or the kernel
was unavailable -> graceful fallback, never a load failure).
diffusers keeps a process-wide active backend that ``set_attention_backend`` also updates, and
a fresh transformer's processors follow it (default None). So a load wanting native must
restore it explicitly, else it inherits a backend an earlier load pinned (e.g. cuDNN under a
speed profile), breaking the ``off`` guarantee. Best-effort."""
setters = [
s
for s in (getattr(t, "set_attention_backend", None) for t in _attention_dits(pipe))
if callable(s)
]
if not setters:
return None
if backend is not None:
_ensure_attention_backend_installed(backend, logger)
engaged = False
for fn in setters:
try:
fn(backend)
engaged = True
except Exception as exc: # noqa: BLE001 — unavailable kernel -> restore native below
_warn(logger, backend, exc)
if engaged:
# set_attention_backend also pins the backend process-wide. Each DiT's processors keep it
# locally, so reset the global to native ONCE, else a later component inherits this kernel.
_reset_global_backend_to_native(logger)
if logger is not None:
logger.info("diffusion.attention: backend=%s", backend)
return backend
# No backend requested, or every set failed: pin native so a stale process-wide backend can't
# leak in. Fresh DiTs follow the global, so one reset covers them all.
_restore_native_backend(setters[0], 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 (AttentionBackendName, fn) or None; read element 0's .value,
# not the tuple (which 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 the process-wide active backend to native after a successful per-transformer set, so
a later unconfigured component doesn't inherit this kernel (the DiT's own processors keep it).
Best-effort: 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)