Repository navigation
feat(attention): add opt-in SageAttention backend for diffusion models - #9657
Pfannkuchensack wants to merge 6 commits into
Conversation
New `attention_backend: auto | sage` setting; with `sage`, eligible PyTorch SDPA calls inside a denoising scope run on SageAttention 2.2 and every other call keeps SDPA unchanged. Each GPU's first SageAttention call is checked against SDPA and a failing GPU falls back for the session; no model family enters the scope yet. Adds the SageAttention docs page and the regenerated config schema.
lstein
left a comment
There was a problem hiding this comment.
Reviewed at c2efc7f. The fallback path is solid: outside the scope every call reaches torch SDPA with its arguments untouched, the per-thread ContextVar scope nests and resets correctly, the first-call GQA slice (query[:1, :group] vs key[:1, :1]) is the right head group, install cannot raise, and ROCm is never touched. Unit tests pass locally (slow tests not run; no SageAttention on this machine).
One blocker, then non-blocking notes.
Blocker
1. The kernel table does not match sageattn() for sm86 (RTX 30-series, A10/A40).
_kernel (invokeai/backend/util/sage_attention.py:134) sends every major == 8, minor != 9 device to sageattn_qk_int8_pv_fp16_cuda(pv_accum_dtype="fp32"), and the docstring says this follows sageattn's own dispatch. It does not: upstream sageattention/core.py dispatches
if arch == "sm80":
return sageattn_qk_int8_pv_fp16_cuda(..., pv_accum_dtype="fp32")
elif arch == "sm86":
return sageattn_qk_int8_pv_fp16_triton(...)and has no sm87 branch at all. So on an RTX 3090 this PR runs a kernel that SageAttention itself deliberately avoids on that architecture. At best it is an unvalidated path (only sm89 was measured); at worst the launch fails, the first-call check retires the GPU, and the feature never works on the most common Ampere consumer cards. TestKernelTable restates the same mapping, so it cannot catch this. Please either follow sageattn's table for sm86 (the Triton kernel; check whether it takes smooth_k the same way, since _unaccepted_arguments would refuse it otherwise), or drop sm86/sm87 from the table until they are measured, and correct the docstring/PR description.
Non-blocking follow-up
2. With this PR alone, attention_backend: sage has no effect but reports itself enabled. Nothing enters sage_attention_scope() yet, so startup imports SageAttention/Triton and logs "SageAttention … enabled", while the "serves … on cuda:N" line and per-generation debug counts shown in the docs' Enabling section can never appear, and the config description points at an empty Supported models list. This resolves itself as the family PRs (#9658 onward) land; just make sure the stack isn't split across a release with only this PR in it.
Advisory
-
Recovery is overstated for sticky CUDA errors. The docs say a failing GPU "uses PyTorch's attention until restart", but a sticky fault (illegal address, misaligned address,
__trap) poisons the context:_disableruns and then the fallbackoriginal(...)(:377) raises the same error. Only non-sticky failures such as "no kernel image" really fall back. Worth one sentence in the docs. -
Transient failures retire a GPU permanently (
:366-371). Onlytorch.OutOfMemoryErroris treated as transient. A Triton module-load OOM (RuntimeError("Triton Error [CUDA]: out of memory")) on a full card, a Triton cache-renamePermissionErrorwhen two Windows GPUs compile their first call concurrently, or a pending async error from an unrelated earlier kernel surfacing at the check's.item()all disable SageAttention on that device until restart. -
Validation is keyed per device, not per kernel instantiation (
:364). A first call at head_dim 128/bf16 validates the device; a later head_dim 64 or fp16 call is never compared, so a launch failure there would surface outside the fallback. -
OOM inside SageAttention propagates instead of falling back (
:366-368). The quantized copies (~1.7 B per Q/K/V element; ~2 GB per call for Wan 14B at 720p per the docs) are not in the working-memory estimators inbackend/util/attention.py, so a generation that fits with SDPA can OOM here. Falling back tooriginalafter an OOM in the Sage path would be safer, and "peak VRAM unchanged" is only measured on smaller models. -
Output layout differs from SDPA. SageAttention returns a contiguous
[B, H, N, D]tensor, while the fused SDPA kernels return[B, H, N, D]with[B, N, H, D]memory order, so callers'out.transpose(1, 2).reshape(...)becomes a full copy per call instead of a view. No caller uses.viewthere, so this is a cost to the speedup rather than a bug.
Notes for the family PRs:
- Inside the scope, SageAttention bypasses
sdpa_kernel/sdpa_policy. For Krea-2 (#9662) that silently defeatsINVOKE_KREA2_SDPA_BACKEND=cudnn, whose purpose is proving which kernel served a run, andforeign_window_entries()cannot detect it. The docs only mentionDIFFUSERS_ATTN_BACKEND. - "Text encoders and VAEs keep PyTorch SDPA" depends on scope placement: ControlNets or IP-adapter encoders run inside the denoise loop would be in scope.
- In legacy single-device mode with
device: cuda:1, the worker never callsset_device, so_cuda_indexreturns None and every call falls back asdevice, visible only at DEBUG.
Test gaps (fast suite): first-call check under GQA (a key[:1, :group] mutant survives), the autocast and is_nested exclusions, and the transient-vs-permanent error classification.
…ent SageAttention failures sm86 now runs the Triton kernel upstream's sageattn picks there, and sm87 gets none; the Windows build's CUDA routing for both is unexplained and unmeasured here. Out-of-memory and Triton cache errors fall back to SDPA for that call (three retire the device), and the first-call check runs per precision and head size. The kernel runs with the tensors' device current, so a single-device install on cuda:1 uses SageAttention too.
An out-of-memory error now moves only the rest of that generation to SDPA, logged once, instead of counting toward retiring the device; out-of-memory is recognised by the shared is_oom_error. Only Triton cache failures in a row retire a device, and a success resets the count.
# Conflicts: # invokeai/frontend/api/openapi.json
|
@lstein Thanks for the thorough review. Addressed in e9aa6c6 and ce561e3; the branch also merges current Blocker (sm86): you're right. The table mirrored the Windows wheel (woct0rdho), which moved sm86/sm87 to Non-blocking:
Family-PR notes:
Test gaps: the GQA head pairing is covered, with a distinct K/V head so a |
# Conflicts: # invokeai/frontend/api/openapi.json # tests/test_config.py
Summary
At high resolutions and for video, most of a diffusion step is attention, and Windows builds of PyTorch have no flash-attention kernel. This adds an opt-in
attention_backend: sagesetting that runs eligible attention inside the denoising model on SageAttention 2.2, an 8-bit quantized attention kernel. On an RTX 4090 under Windows those calls run 2.0–2.9× faster than PyTorch's fastest SDPA kernel, and whole generations get 3–13 % faster at 1024², 15–35 % at 2048² and 12–35 % for the two measured video models. Images are not pixel-identical toautoat the same seed; in a blind review of 48 image pairs none was judged worse. The default,auto, leaves PyTorch's attention untouched.This PR adds the mechanism, the setting and the docs; no model family uses SageAttention yet. Each family enters the scope in its own stacked PR, with its own measurements:
How it works:
attention_backend: sage,apply_monkeypatchescallsinstall_sage_attention(), which wrapstorch.nn.functional.scaled_dot_product_attentiononce, the way the ROCm SDPA guard already does. That one wrapper covers direct SDPA calls, diffusers attention processors and diffusers'nativedispatch backend. diffusers' ownsagebackend is not used: it raises on any mask, applies process-wide and misses the direct calls.sage_attention_scope(), a per-threadContextVarscope that denoise invocations enter. Text encoders, VAEs, other threads and every call outside the scope keep SDPA as before.device: cuda:1never makes it so). Every other call reaches the original SDPA with its arguments untouched. Withlog_level: debug, each scope logs how many calls ran on SageAttention and why the others did not.sageattn()'s table (sm80: FP16 PV in CUDA; sm86: FP16 PV in Triton; sm87: none; sm89: FP8 PV; sm90: the sm90 FP8 kernel; sm100/120/121: per-warp FP8 PV), but with K smoothing off. The Windows wheel the docs install moved sm86/sm87 to the CUDA kernel without a stated reason; neither is measured here, so upstream's choice stands, and the hardware tests also run the sm86 Triton kernel on the 4090. The default smoothing breaks block 0 of Qwen-Image and Krea-2, where one Q/K channel sits near 600 for all image tokens: cosine against fp32 drops to 0.18 and 0.53, against 0.997 and 0.988 without smoothing, which is also 2–11 % faster. A real Qwen-Image head from that block is a 192 KB test fixture.is_oom_error) moves the rest of that generation to SDPA, which needs less, with one INFO line; the next generation tries again. A Triton cache write failure (two GPUs compiling one kernel at once on Windows) falls back for that call; three in a row retire the GPU.smooth_k), startup logs one warning with the reason and nothing is wrapped.Users install SageAttention themselves;
pyproject.tomlanduv.lockare unchanged. The new docs page covers Windows (triton-windows and a prebuilt wheel, installed into InvokeAI's.venvwithuv pip) and Linux (built from source).Related Issues / Discussions
Related: #9529, which asks for flash-attn or SageAttention to speed up Wan 2.2. The Wan PR #9665 closes it; flash-attn is not part of this stack.
QA Instructions
Checks on this branch:
uv run --no-sync pytest tests/backend/util/test_sage_attention.py tests/test_config.py tests/test_imports.py: 118 passed. The unit tests use a fake kernel that declares the real keyword signature.uv run --no-sync pytest -m slow tests/backend/util/test_sage_attention.pyon an RTX 4090 with SageAttention 2.2.0: 16 passed, each test once with the GPU's own kernel and once with the sm86 Triton kernel. They cover accuracy on FLUX, SDXL, Krea-2 GQA and a partial-tile cross-attention shape, each asserting the kernel served the call; the Qwen-Image outlier head (relative L2 0.10, against 1.62 with K smoothing); results identical to calling the kernel directly inside the scope and to SDPA outside it and for masked or autocast calls; and a first call that needs no more memory than later ones.ruff checkandruff format --check;pnpm -C docs test,pnpm -C docs build(all internal links valid) andpnpm -C docs check-deploy-output.openapi.json,schema.tsandsettings.jsonwere regenerated; the new field is the only change.Manual:
auto, FLUX.1 at 1024² and a fixed seed is bit-identical tomain, and no SageAttention line is logged.attention_backend: sageand restart. Startup logsSageAttention 2.2.0+cu130torch2.10.0andhigher.post6 enabled for the attention inside supported diffusion models .... Without the package, startup logs one warning that names the reason.Measurements behind the stack, on an RTX 4090 with Windows 10, torch 2.13.0+cu130, triton-windows 3.7.1.post27 and SageAttention 2.2.0 (post6 Windows wheel):
Whole generation, model loaded:
sageattn()'s default K smoothing. The selection shipped here is 2–11 % faster per call; for one FLUX.1 seed at 1024², its LPIPS against SDPA dropped from 0.128 to 0.041.Review
Resolved:
sageattn()'s CUDA rule (fp32+fp16 only from CUDA 12.8), Blackwell is limited to the capabilities SageAttention has kernels for, installation can no longer raise, and a kernel whose signature would silently dropsmooth_kis refused.device: cuda:1in single-device mode, every call fell back silently because the worker never makes that device current; the wrapper now does for the call.Remaining risks:
autoat the same seed. In Qwen-Image block 52, V reaches about 3500 in every channel, which costs every SageAttention variant about 17 % relative error per call (bf16 SDPA: 0.2 %); those images were still judged equivalent.[B, H, N, D]tensor where the fused SDPA kernels return[B, N, H, D]memory order, so callers'transpose(1, 2).reshape(...)copies. The measured end-to-end times include that copy.Compatibility / Rollout
attention_backend(autoorsage, defaultauto, environment variableINVOKEAI_ATTENTION_BACKEND). No migration; withautonothing is wrapped.openapi.json,schema.tsanddocs/src/generated/settings.jsonadd only this field; the frontend does not use it.mainand shows the earlier commits until they merge. This PR alone enables nothing (the startup line says SageAttention is enabled, but no model enters its scope), so it should not ship in a release without at least one family PR.attention_backend: autoor uninstall the package.Checklist
What's Newcopy (if doing a release after this PR)