Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
c2efc7f
feat(attention): add opt-in SageAttention backend for diffusion models
Pfannkuchensack Oct 3, 2026
e864d62
feat(flux): run FLUX.1 attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
a1e97b7
feat(flux2): run FLUX.2 attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
6c87189
feat(z-image): run Z-Image attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
4abdcd4
feat(qwen-image): run Qwen-Image attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
482e7f3
feat(krea2): run Krea-2 attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
1968eb9
feat(sdxl): run SDXL attention on SageAttention when enabled
Pfannkuchensack Oct 3, 2026
e9aa6c6
fix(attention): follow upstream's sm86 kernel and recover from transi…
Pfannkuchensack Oct 9, 2026
ce561e3
fix(attention): keep a GPU on SageAttention after running out of memory
Pfannkuchensack Oct 9, 2026
7e4f914
Merge remote-tracking branch 'upstream/main' into feat/sage-attention
Pfannkuchensack Oct 9, 2026
29fd425
Merge branch 'feat/sage-attention' into feat/sage-attention-flux1
Pfannkuchensack Oct 9, 2026
bac9e01
Merge branch 'feat/sage-attention-flux1' into feat/sage-attention-flux2
Pfannkuchensack Oct 9, 2026
65cdc4b
Merge branch 'feat/sage-attention-flux2' into feat/sage-attention-z-i…
Pfannkuchensack Oct 9, 2026
f320473
Merge branch 'feat/sage-attention-z-image' into feat/sage-attention-q…
Pfannkuchensack Oct 9, 2026
a5e80a0
Merge branch 'feat/sage-attention-qwen-image' into feat/sage-attentio…
Pfannkuchensack Oct 9, 2026
ec98c09
Merge branch 'feat/sage-attention-krea2' into feat/sage-attention-sdxl
Pfannkuchensack Oct 9, 2026
4a41e47
Merge branch 'main' into feat/sage-attention
Pfannkuchensack Oct 9, 2026
971b661
Merge remote-tracking branch 'upstream/main' into feat/sage-attention
Pfannkuchensack Oct 10, 2026
a599781
Merge branch 'feat/sage-attention' into feat/sage-attention-flux1
Pfannkuchensack Oct 10, 2026
c16c7f5
Merge branch 'feat/sage-attention-flux1' into feat/sage-attention-flux2
Pfannkuchensack Oct 10, 2026
8434937
Merge branch 'feat/sage-attention-flux2' into feat/sage-attention-z-i…
Pfannkuchensack Oct 10, 2026
4a9b290
Merge branch 'feat/sage-attention-z-image' into feat/sage-attention-q…
Pfannkuchensack Oct 10, 2026
f50c8fa
Merge branch 'feat/sage-attention-qwen-image' into feat/sage-attentio…
Pfannkuchensack Oct 10, 2026
06b931a
Merge branch 'feat/sage-attention-krea2' into feat/sage-attention-sdxl
Pfannkuchensack Oct 10, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
148 changes: 148 additions & 0 deletions docs/src/content/docs/configuration/Optimization/sage-attention.mdx
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
---
title: SageAttention
sidebar:
order: 5
---

import { Steps } from '@astrojs/starlight/components';

[SageAttention](https://github.com/thu-ml/SageAttention) is a faster attention kernel for NVIDIA GPUs. It quantizes the
attention inputs to 8 bits inside the kernel. On an RTX 4090 under Windows that makes the attention of the image and
video models two to three times faster than PyTorch's own kernels. How much of a generation that saves depends on how
much of each step is attention: it grows with resolution and video length.

It is **opt-in**. Because the kernel is quantized, an image made with SageAttention is not pixel-identical to one made
without it at the same seed — composition details can shift the way they do between two GPU models. In a blind
side-by-side comparison of 48 image pairs from six image model families, no pair was judged worse with SageAttention;
the video models were checked against PyTorch's attention for picture changes and frame-to-frame flicker.

## Supported models

SageAttention is enabled per model family once that family has been measured with it. All other models keep PyTorch's
attention.

Whole generation (text encoding, denoising and decoding) with the model already loaded, on an RTX 4090 under Windows,
compared with PyTorch's default attention there:

| Model (as measured) | 1024×1024 | 2048×2048 | Notes |
| --- | --- | --- | --- |
| FLUX.1 dev (FP8) | 16.1 s → 14.0 s (−13 %) | 84.2 s → 54.8 s (−35 %) | Also serves ControlNets inside the denoising loop. |
| FLUX.2 Klein 9B (FP8) | 15.6 s → 15.1 s (−3 %) | 27.1 s → 23.1 s (−15 %) | |
| Z-Image (FP8) | 10.4 s → 10.0 s (−5 %) | 33.7 s → 25.9 s (−23 %) | Z-Image Control masks its attention and keeps PyTorch's. |
| Qwen-Image 2512 (GGUF Q4_K_M) | 78.8 s → 73.5 s (−7 %) | 208.9 s → 150.3 s (−28 %) | One block reaches about 17 % error per call; images were still judged equivalent. |
| Krea-2 Turbo (GGUF Q4_K_M) | 11.5 s → 11.0 s (−4 %) | 42.7 s → 36.2 s (−15 %) | Off while `INVOKE_KREA2_SDPA_BACKEND` pins a PyTorch kernel. |
| SDXL | 4.4 s → 4.2 s (−5 %) | 19.6 s → 16.1 s (−18 %) | Text cross-attention stays on PyTorch (too short to gain). Also serves ControlNets. |

At 1024×1024 most of a step is spent outside attention, so the gain is small; at high resolutions and for video,
attention dominates. GGUF models spend much of each step dequantizing weights, which SageAttention does not touch.
Windows' PyTorch has no flash-attention kernel; on Linux PyTorch's own baseline is faster and the gain correspondingly
smaller (not measured).

## Requirements

- **An NVIDIA GPU with compute capability 8.0 or newer** — RTX 30-series, RTX 40-series, RTX 50-series, A100, H100.
InvokeAI has measured it on an RTX 4090; the other generations use the kernels SageAttention picks for them, which
InvokeAI has not validated (on RTX 30-series cards, its Triton kernel). Jetson Orin is not supported.
- **SageAttention 2.2, installed separately.** InvokeAI was validated with 2.2.0. The `sageattention` package on PyPI is
the older SageAttention 1 and is not used. On Windows it also needs the `triton-windows` package.
- Not available on AMD (ROCm), Apple Silicon (MPS), Intel (XPU) or CPU. There InvokeAI keeps PyTorch's attention.

## Installing

SageAttention goes into the Python environment InvokeAI runs from — the `.venv` folder in your install directory. The
commands use [uv](https://docs.astral.sh/uv/), which InvokeAI's environment is managed with; install it first if
`uv --version` does not work in your terminal.

### Windows

<Steps>
1. Open a terminal in your InvokeAI install directory.
2. Install Triton for Windows, matching the Triton version of InvokeAI's PyTorch (Triton 3.7 for PyTorch 2.13):

```powershell
uv pip install --python .venv\Scripts\python.exe "triton-windows>=3.7,<3.8"
```

3. Install a SageAttention 2.2 wheel from the [SageAttention for Windows releases](https://github.com/woct0rdho/SageAttention/releases)
that matches InvokeAI's PyTorch (CUDA 13, PyTorch 2.10 or newer), without letting it touch PyTorch:

```powershell
uv pip install --python .venv\Scripts\python.exe --no-deps "https://github.com/woct0rdho/SageAttention/releases/download/v2.2.0-windows.post6/sageattention-2.2.0%2Bcu130torch2.10.0andhigher.post6-cp310-abi3-win_amd64.whl"
```
</Steps>

### Linux

SageAttention 2 has no prebuilt wheels for Linux. Build it from source into InvokeAI's environment
(`.venv/bin/python`) as described in [SageAttention's installation instructions](https://github.com/thu-ml/SageAttention#installation),
with the CUDA toolkit that matches InvokeAI's PyTorch (CUDA 13) installed. InvokeAI's environment has no `pip`; use
`uv pip install --python .venv/bin/python ...` where those instructions say `pip install`. Triton ships with PyTorch on
Linux.

:::caution[Keep it matched to InvokeAI's PyTorch]
SageAttention is built against PyTorch, and Triton must match PyTorch's version. When an update moves InvokeAI to a new
PyTorch, install the matching `triton-windows` again on Windows, and on Linux rebuild SageAttention. A wheel for a
different CUDA major version, or a repair that re-creates the environment, needs a matching SageAttention too. Until
then InvokeAI logs why it cannot use it and generates with PyTorch's attention.
:::

## Enabling

Set the backend in your [`invokeai.yaml`](/configuration/invokeai-yaml/) and restart InvokeAI:

```yaml
attention_backend: sage
```

The startup log confirms it:

```
SageAttention 2.2.0+cu130torch2.10.0andhigher.post6 enabled for the attention inside supported diffusion models ...
```

and the first generation that uses it names the GPU and the kernel:

```
SageAttention 2.2.0+cu130torch2.10.0andhigher.post6 serves diffusion-model attention on cuda:0 (NVIDIA GeForce RTX 4090, sm_89, sageattn_qk_int8_pv_fp8_cuda).
```

If SageAttention cannot be used, the startup log says why — package missing, Triton missing, GPU too old — and
InvokeAI generates with PyTorch's attention as usual. With `log_level: debug`, every generation logs how many attention
calls ran on SageAttention and why the others did not. To go back, set `attention_backend: auto`.

## What stays on PyTorch's attention

SageAttention serves the attention inside the denoising model of the families under [Supported models](#supported-models).
Everything else keeps PyTorch's attention:

- **Text encoders, image encoders and VAEs.** They run outside the denoising loop. What runs inside it, such as a
ControlNet, is served like the model itself.
- **Masked attention**, such as regional prompts. SageAttention has no mask input.
- **Short sequences** — fewer than 1536 query or 512 key tokens. There the cost of quantizing the inputs outweighs the
faster kernel.
- **Head sizes other than 64 and 128.**
- **SDXL with `attention_type: sliced`** — sliced attention does not go through PyTorch's attention kernel.
- **Models not listed under Supported models**, such as SD 1.5 and SD 2.

Pinning a PyTorch kernel through diffusers' `DIFFUSERS_ATTN_BACKEND`, or in a custom node with
`torch.nn.attention.sdpa_kernel`, does not stop SageAttention; set `attention_backend: auto` to measure PyTorch's own
kernels.

## Notes

- **The first generation pauses briefly.** SageAttention compiles a few Triton kernels on first use, about two seconds
on an RTX 4090; they are cached on disk afterwards.
- **Each GPU's first SageAttention call is checked against PyTorch's attention**, once for each precision and head
size. If the result is wrong or the kernel fails, InvokeAI logs a warning once and that GPU uses PyTorch's attention
until restart. Some GPU faults (an illegal memory access, for example) leave the GPU unusable for the rest of the
process: then PyTorch's attention fails the same way, and only a restart helps. A fault the GPU raises later, while
running a kernel, cannot be caught at all and ends the generation.
- **Memory.** Peak VRAM was unchanged for every measured model. Each call holds quantized copies of its inputs, about
1.7 bytes per element of the query, key and value together — around 70 MB for FLUX.1 at 1024×1024, and around 2 GB
per call for Wan 2.2 14B at 1280×720 with 81 frames (not measured). When a call does not fit, the rest of that
generation runs on PyTorch's attention, which needs less, and the log says so; the next generation tries
SageAttention again.
- **Recall.** `attention_backend` is not stored in an image's metadata, so an image recalled after switching it is
generated slightly differently.
- **Black or noisy images** have been reported with SageAttention for some models in other applications. None appeared
in InvokeAI's tests; if you see one, set `attention_backend: auto` and report it with the model and settings.
2 changes: 1 addition & 1 deletion docs/src/content/docs/troubleshooting/faq.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ Follow the same steps to scan and import the missing models.

## Triton error on startup

This can be safely ignored. Invoke doesn't use Triton, but if you are on Linux and wish to dismiss the error, you can install Triton.
This can be safely ignored. Invoke only uses Triton for [SageAttention](/configuration/optimization/sage-attention/), which is off unless you enable it. If you are on Linux and wish to dismiss the error, you can install Triton.

## Unable to Copy on Firefox

Expand Down
14 changes: 14 additions & 0 deletions docs/src/generated/settings.json
Original file line number Diff line number Diff line change
Expand Up @@ -798,6 +798,20 @@
"type": "typing.Literal['auto', 'balanced', 'max', 1, 2, 3, 4, 5, 6, 7, 8]",
"validation": {}
},
{
"category": "GENERATION",
"default": "auto",
"description": "Attention kernel for the attention inside diffusion models. `auto` lets PyTorch choose among its exact SDPA kernels and never selects a quantized one. `sage` uses SageAttention 2 for the models listed in the SageAttention docs: faster, most at high resolutions and for video, but quantized, so images differ slightly from `auto` at the same seed. SageAttention is installed separately and needs an NVIDIA GPU with compute capability 8.0 or newer; masked attention, text encoders and VAEs keep PyTorch SDPA.",
"env_var": "INVOKEAI_ATTENTION_BACKEND",
"literal_values": [
"auto",
"sage"
],
"name": "attention_backend",
"required": false,
"type": "typing.Literal['auto', 'sage']",
"validation": {}
},
{
"category": "GENERATION",
"default": false,
Expand Down
42 changes: 22 additions & 20 deletions invokeai/app/invocations/flux/flux_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@
from invokeai.backend.stable_diffusion.diffusers_pipeline import PipelineIntermediateState
from invokeai.backend.stable_diffusion.diffusion.conditioning_data import FLUXConditioningInfo
from invokeai.backend.util.devices import TorchDevice
from invokeai.backend.util.sage_attention import sage_attention_scope


@invocation(
Expand Down Expand Up @@ -525,26 +526,27 @@ def _run_diffusion(
else:
context.logger.debug(f"DyPE disabled: resolution={self.width}x{self.height}, preset={self.dype_preset}")

x = denoise(
model=transformer,
img=x,
img_ids=img_ids,
pos_regional_prompting_extension=pos_regional_prompting_extension,
neg_regional_prompting_extension=neg_regional_prompting_extension,
timesteps=timesteps,
step_callback=self._build_step_callback(context),
guidance=self.guidance,
cfg_scale=cfg_scale,
inpaint_extension=inpaint_extension,
controlnet_extensions=controlnet_extensions,
pos_ip_adapter_extensions=pos_ip_adapter_extensions,
neg_ip_adapter_extensions=neg_ip_adapter_extensions,
img_cond=img_cond,
img_cond_seq=img_cond_seq,
img_cond_seq_ids=img_cond_seq_ids,
dype_extension=dype_extension,
scheduler=scheduler,
)
with sage_attention_scope():
x = denoise(
model=transformer,
img=x,
img_ids=img_ids,
pos_regional_prompting_extension=pos_regional_prompting_extension,
neg_regional_prompting_extension=neg_regional_prompting_extension,
timesteps=timesteps,
step_callback=self._build_step_callback(context),
guidance=self.guidance,
cfg_scale=cfg_scale,
inpaint_extension=inpaint_extension,
controlnet_extensions=controlnet_extensions,
pos_ip_adapter_extensions=pos_ip_adapter_extensions,
neg_ip_adapter_extensions=neg_ip_adapter_extensions,
img_cond=img_cond,
img_cond_seq=img_cond_seq,
img_cond_seq_ids=img_cond_seq_ids,
dype_extension=dype_extension,
scheduler=scheduler,
)

x = unpack(x.float(), self.height, self.width)
return x
Expand Down
40 changes: 21 additions & 19 deletions invokeai/app/invocations/flux2/flux2_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@
from invokeai.backend.stable_diffusion.diffusion.conditioning_data import FLUXConditioningInfo
from invokeai.backend.util.attention import sdpa_score_matrix_bytes
from invokeai.backend.util.devices import TorchDevice
from invokeai.backend.util.sage_attention import sage_attention_scope

# FLUX.2 attention geometry. The head dim is 128 across every variant and the head count follows
# the hidden size (Klein 4B: 3072/24, Klein 9B: 4096/32, [dev] 6144/48), so the width is the single
Expand Down Expand Up @@ -623,25 +624,26 @@ def _run_diffusion(self, context: InvocationContext) -> torch.Tensor:
"Regional masks will be ignored for this generation."
)

x = denoise(
model=transformer,
img=x,
img_ids=img_ids,
txt=txt,
txt_ids=txt_ids,
timesteps=timesteps,
step_callback=self._build_step_callback(context),
guidance=self.guidance,
cfg_scale=cfg_scale_list,
neg_txt=neg_txt,
neg_txt_ids=neg_txt_ids,
scheduler=scheduler,
mu=mu,
inpaint_extension=inpaint_extension,
img_cond_seq=img_cond_seq,
img_cond_seq_ids=img_cond_seq_ids,
pos_joint_attention_kwargs=pos_joint_attention_kwargs,
)
with sage_attention_scope():
x = denoise(
model=transformer,
img=x,
img_ids=img_ids,
txt=txt,
txt_ids=txt_ids,
timesteps=timesteps,
step_callback=self._build_step_callback(context),
guidance=self.guidance,
cfg_scale=cfg_scale_list,
neg_txt=neg_txt,
neg_txt_ids=neg_txt_ids,
scheduler=scheduler,
mu=mu,
inpaint_extension=inpaint_extension,
img_cond_seq=img_cond_seq,
img_cond_seq_ids=img_cond_seq_ids,
pos_joint_attention_kwargs=pos_joint_attention_kwargs,
)

# Apply BN denormalization if BN stats are available
# The diffusers Flux2KleinPipeline applies: latents = latents * bn_std + bn_mean
Expand Down
6 changes: 6 additions & 0 deletions invokeai/app/invocations/krea2/krea2_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@
from invokeai.backend.util.attention import sdpa_score_matrix_bytes
from invokeai.backend.util.devices import TorchDevice
from invokeai.backend.util.logging import InvokeAILogger
from invokeai.backend.util.sage_attention import sage_attention_scope

# Krea-2 latent channels (Qwen-Image VAE z_dim). The packed transformer in_channels is 16 * patch_size**2 = 64.
KREA2_LATENT_CHANNELS = 16
Expand Down Expand Up @@ -625,6 +626,11 @@ def _run_diffusion(self, context: InvocationContext):
# the sampler will use.
style_extension.prepare([sigma.item() for sigma in sigmas_sched[:total_steps]])

# An INVOKE_KREA2_SDPA_BACKEND override pins one PyTorch kernel to measure it; SageAttention must not
# serve those calls instead.
if resolve_krea2_sdpa_backends().override is None:
exit_stack.enter_context(sage_attention_scope())

benchmark = _Krea2StepBenchmark.create(device)

for step_idx, t in enumerate(tqdm(timesteps_sched)):
Expand Down
3 changes: 3 additions & 0 deletions invokeai/app/invocations/qwen_image/qwen_image_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
from invokeai.backend.stable_diffusion.diffusers_pipeline import PipelineIntermediateState
from invokeai.backend.stable_diffusion.diffusion.conditioning_data import QwenImageConditioningInfo
from invokeai.backend.util.devices import TorchDevice
from invokeai.backend.util.sage_attention import sage_attention_scope


@invocation(
Expand Down Expand Up @@ -485,6 +486,8 @@ def _run_diffusion(self, context: InvocationContext):
)
)

exit_stack.enter_context(sage_attention_scope())

for step_idx, t in enumerate(tqdm(timesteps_sched)):
# The pipeline passes timestep / 1000 to the transformer
timestep = t.expand(latents.shape[0]).to(inference_dtype)
Expand Down
Loading
Loading