Skip to content

feat(attention): add opt-in SageAttention backend for diffusion models - #9657

Open
Pfannkuchensack wants to merge 6 commits into
invoke-ai:mainfrom
Pfannkuchensack:feat/sage-attention
Open

Pfannkuchensack wants to merge 6 commits into
invoke-ai:mainfrom
Pfannkuchensack:feat/sage-attention

Conversation

@Pfannkuchensack

@Pfannkuchensack Pfannkuchensack commented Oct 3, 2026 •

Copy link
Copy Markdown
Member

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: sage setting 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 to auto at 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:

PR Enables Measured with
#9658 FLUX.1 FLUX.1 dev FP8
#9659 FLUX.2 FLUX.2 Klein 9B FP8
#9660 Z-Image Z-Image FP8
#9661 Qwen-Image Qwen-Image 2512 GGUF Q4_K_M
#9662 Krea-2 Krea-2 Turbo GGUF Q4_K_M
#9663 SDXL SDXL
#9664 LTX-2 LTX-2.5 distilled int8
#9665 Wan 2.2 Wan 2.2 TI2V-5B GGUF Q8

How it works:

  • With attention_backend: sage, apply_monkeypatches calls install_sage_attention(), which wraps torch.nn.functional.scaled_dot_product_attention once, the way the ROCm SDPA guard already does. That one wrapper covers direct SDPA calls, diffusers attention processors and diffusers' native dispatch backend. diffusers' own sage backend is not used: it raises on any mask, applies process-wide and misses the direct calls.
  • The wrapper acts only inside sage_attention_scope(), a per-thread ContextVar scope that denoise invocations enter. Text encoders, VAEs, other threads and every call outside the scope keep SDPA as before.
  • Inside the scope a call goes to SageAttention only if it has no mask, dropout or causal flag; is 4-D fp16/bf16 with head dim 64 or 128; has at least 1536 query and 512 key tokens (the measured break-even: below it, quantizing costs more than the kernel saves); has divisible GQA heads; needs no grad, runs without autocast and has q, k and v on one CUDA device, which is made current for the call (SageAttention launches on the current device; a single-device install with device: cuda:1 never makes it so). Every other call reaches the original SDPA with its arguments untouched. With log_level: debug, each scope logs how many calls ran on SageAttention and why the others did not.
  • The kernel is chosen per compute capability, following upstream 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.
  • The first SageAttention call per GPU, precision and head size is compared with SDPA on a slice (the last 1024 query rows of the first sample and K/V group; relative L2 at most 0.5). A wrong result or a kernel error moves that GPU back to SDPA with one warning.
  • Out of memory (recognised by the shared 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.
  • If SageAttention cannot be used (package missing, SageAttention 1 from PyPI, Triton missing, ROCm, no CUDA GPU of compute capability 8.0 or newer, or a kernel signature that would drop smooth_k), startup logs one warning with the reason and nothing is wrapped.

Users install SageAttention themselves; pyproject.toml and uv.lock are unchanged. The new docs page covers Windows (triton-windows and a prebuilt wheel, installed into InvokeAI's .venv with uv 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.py on 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 check and ruff format --check; pnpm -C docs test, pnpm -C docs build (all internal links valid) and pnpm -C docs check-deploy-output. openapi.json, schema.ts and settings.json were regenerated; the new field is the only change.

Manual:

  1. With the default auto, FLUX.1 at 1024² and a fixed seed is bit-identical to main, and no SageAttention line is logged.
  2. Install SageAttention as the new docs page describes, set attention_backend: sage and restart. Startup logs SageAttention 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.
  3. With this PR alone no generation uses SageAttention; the stacked PRs enable it per family.

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

  • Each seed ran three arms in alternating order: PyTorch's default SDPA, SDPA with cuDNN first (exact, so it shows how far any kernel change moves an image) and SageAttention.
  • Step time is the slope between two step counts over three repeats.
  • Quality is LPIPS against the default SDPA at the same seed (3 prompts × 4 seeds at 1024², 2 × 2 at 2048²), plus a blind side-by-side review of 48 image pairs.

Whole generation, model loaded:

Model 1024² 2048² Blind review (8 pairs)
FLUX.1 dev FP8 16.1 → 14.0 s (−13 %) 84.2 → 54.8 s (−35 %) 8 equal
FLUX.2 Klein 9B FP8 15.6 → 15.1 s (−3 %) 27.1 → 23.1 s (−15 %) 8 equal
Z-Image FP8 10.4 → 10.0 s (−5 %) 33.7 → 25.9 s (−23 %) 8 equal
Qwen-Image 2512 GGUF Q4_K_M 78.8 → 73.5 s (−7 %) 208.9 → 150.3 s (−28 %) 8 equal
Krea-2 Turbo GGUF Q4_K_M 11.5 → 11.0 s (−4 %) 42.7 → 36.2 s (−15 %) 8 equal
SDXL 4.4 → 4.2 s (−5 %) 19.6 → 16.1 s (−18 %) 7 equal, 1 better
Video model Size Whole generation Per-frame LPIPS (cuDNN-first) Flicker ratio
LTX-2.5 distilled int8 1248×704, 121 frames 72.0 → 63.4 s (−12 %) 0.142 (0.090) 1.001
Wan 2.2 TI2V-5B GGUF Q8 1280×704, 121 frames 380.9 → 247.9 s (−35 %) 0.104 (0.079) 0.998
  • Peak VRAM inside the denoise scope was the same in all three arms for every model.
  • The wrapper costs 0.4 µs per SDPA call outside the scope and 2.1 µs per call that falls back inside it.
  • The timings, LPIPS and blind review come from the measurement prototype, which still used 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.
  • Not verified: Linux, GPUs other than sm89 (Ampere, Hopper and Blackwell use SageAttention's kernels for them), multi-GPU.

Review

Resolved:

  • The first-call check originally recomputed the whole call against SDPA, at about 14 bytes per element, roughly 5 GB for Wan 14B at 720p. It now checks a fixed slice. On Wan 5B at 1280×704 with 121 frames, the first call measured 1604 MiB extra, the same as every later call.
  • Accumulation follows 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 drop smooth_k is refused.
  • sm86 ran the CUDA kernel upstream avoids there, and sm87 one upstream has none for; both now follow upstream.
  • An out-of-memory error used to propagate, and any other failure retired the GPU for good, including Triton cache races; see the two bullets above. The first-call check ran once per GPU, so a later precision or head size went unchecked; it now runs per variant.
  • With device: cuda:1 in 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:

  • Attention is quantized, so images differ from auto at 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.
  • K smoothing off is right for every measured family, but a synthetic Q/K outlier pattern breaks it. That is why each family is enabled only after it has been measured.
  • A GPU fault raised asynchronously after the checked first call cannot be caught and ends the generation. A fault that poisons the CUDA context (an illegal address, a trap) retires the GPU, but PyTorch's SDPA then fails the same way; only a restart helps. The docs say so.
  • SageAttention returns a contiguous [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

  • New config field attention_backend (auto or sage, default auto, environment variable INVOKEAI_ATTENTION_BACKEND). No migration; with auto nothing is wrapped.
  • openapi.json, schema.ts and docs/src/generated/settings.json add only this field; the frontend does not use it.
  • Merge order: this PR, then the family PRs in the order above. Each family PR targets main and 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.
  • To back out at runtime, set attention_backend: auto or uninstall the package.

Checklist

  • The PR has a short but descriptive title, suitable for a changelog
  • Meaningful regression coverage added / updated where needed; obsolete tests/code removed
  • Persisted-state and API changes include required migrations / compatibility validation
  • Relevant performance/efficiency opportunities considered; material claims have evidence
  • Material review findings resolved and relevant checks rerun
  • Documentation added / updated (if applicable)
  • Updated What's New copy (if doing a release after this PR)

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 lstein left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

  1. 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: _disable runs and then the fallback original(...) (:377) raises the same error. Only non-sticky failures such as "no kernel image" really fall back. Worth one sentence in the docs.

  2. Transient failures retire a GPU permanently (:366-371). Only torch.OutOfMemoryError is treated as transient. A Triton module-load OOM (RuntimeError("Triton Error [CUDA]: out of memory")) on a full card, a Triton cache-rename PermissionError when 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.

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

  4. 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 in backend/util/attention.py, so a generation that fits with SDPA can OOM here. Falling back to original after an OOM in the Sage path would be safer, and "peak VRAM unchanged" is only measured on smaller models.

  5. 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 .view there, 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 defeats INVOKE_KREA2_SDPA_BACKEND=cudnn, whose purpose is proving which kernel served a run, and foreign_window_entries() cannot detect it. The docs only mention DIFFUSERS_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 calls set_device, so _cuda_index returns None and every call falls back as device, 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
@Pfannkuchensack

Copy link
Copy Markdown
Member Author

@lstein Thanks for the thorough review. Addressed in e9aa6c6 and ce561e3; the branch also merges current main, and the eight family PRs are updated on top.

Blocker (sm86): you're right. The table mirrored the Windows wheel (woct0rdho), which moved sm86/sm87 to sageattn_qk_int8_pv_fp16_cuda without a stated reason, not upstream. It now follows upstream: Triton kernel on sm86 (it takes smooth_k under that name, so _unaccepted_arguments passes), nothing on sm87, and the docstring/docs say where the two differ. With no Ampere card here, the hardware test class now runs every test twice: once with the GPU's own kernel and once with the sm86 Triton kernel driven through the wrapper (it runs on the 4090). That includes the Qwen-Image outlier fixture and the memory bound. All 16 pass.

Non-blocking:

  1. This PR alone: agreed. The PR description now says it should not ship in a release without at least one family PR.
  2. Sticky faults: the docs now say that a context-poisoning fault makes SDPA fail the same way and only a restart helps.
  3. Transient failures: out of memory (now via the shared is_oom_error, so Triton's RuntimeError OOM counts) moves the rest of that sage_attention_scope to SDPA with one INFO line, and the next generation tries again. That avoids both retiring the device and flushing the allocator on every call. A Triton cache OSError falls back for that call; three in a row retire the device, and a success resets the count. A pending async error from an unrelated kernel still retires, the conservative choice.
  4. Validation per instantiation: the first-call check is now keyed by (device, dtype, head_dim).
  5. OOM: falls back instead of propagating, as in 4.
  6. Output layout: true, the copy remains. The end-to-end timings in the description were measured with it, so the gains are net of it. I left it as is.

Family-PR notes:

  • Krea-2: feat(krea2): run Krea-2 attention on SageAttention when enabled #9662 already leaves the scope closed while INVOKE_KREA2_SDPA_BACKEND is set. The docs now also say that torch.nn.attention.sdpa_kernel in a custom node does not stop SageAttention.
  • ControlNets etc.: the docs now say that what runs inside the denoise loop, such as a ControlNet, is served like the model.
  • Legacy device: cuda:1: _cuda_index no longer requires the current device. The kernel and the check run under torch.cuda.device(index), so that setup uses SageAttention instead of falling back silently.

Test gaps: the GQA head pairing is covered, with a distinct K/V head so a key[:1, :group] mutant fails, as are autocast, nested queries, the per-variant check, the OOM, cache and permanent-error classes, and the device context.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

7.0.0 backend PRs that change backend files docs PRs that change docs python PRs that change python files python-tests PRs that change python tests services PRs that change app services

Projects

Status: 7.0 Theme: Modular Design

Development

Successfully merging this pull request may close these issues.

3 participants