Buckets:

|
download
raw
12.1 kB

Attention backends

The attention dispatcher is an experimental feature. Please open an issue if you have any feedback or encounter any problems.

Most Diffusers transformer models route attention through an attention dispatcher so you can switch optimized backends behind one API. The dispatcher manages registered implementations and exposes a unified call path for them. Some models, such as many autoencoders, don't use the dispatcher because of their internals, and switching the backend has no effect on their attention layers.

Refer to the table below for an overview of the available attention families and to the Available backends section for a more complete list. The fastest backend depends on the model, GPU, dtype, and input shape.

attention family main feature
FlashAttention minimizes memory reads/writes through tiling and recomputation
AI Tensor Engine for ROCm FlashAttention implementation optimized for AMD ROCm accelerators
SageAttention quantizes attention to int8
FlexAttention PyTorch FlexAttention
PyTorch native built-in PyTorch implementation using scaled_dot_product_attention
xFormers memory-efficient attention with support for various attention kernels

Hub backends (the *_hub names) only need the Kernels library and download the kernel on first use. Other backends need their own package, such as flash-attn or sageattention. The Available backends table lists the requirements Diffusers checks when you enable a backend.

Set a backend on the model

The set_attention_backend() method walks the model’s attention layers and applies the chosen backend on each one. It also sets the dispatcher’s process-wide active backend to the same value.

reset_attention_backend() clears the backend on attention layers only. It does not clear the process-wide active backend. For a temporary switch that restores the previous active backend on exit, use the attention_backend context manager.

The example below enables _flash_3_hub (FlashAttention-3 from the Hub) with device_map="cuda".

import torch
from diffusers import QwenImagePipeline

pipeline = QwenImagePipeline.from_pretrained(
    "Qwen/Qwen-Image", dtype=torch.bfloat16, device_map="cuda"
)
pipeline.transformer.set_attention_backend("_flash_3_hub")

prompt = """
cinematic film still of a cat sipping a margarita in a pool in Palm Springs, California
highly detailed, high budget hollywood movie, cinemascope, moody, epic, gorgeous, film grain
"""
pipeline(prompt).images[0]

The non-Hub FlashAttention-3 backends (_flash_3, _flash_varlen_3) require building FlashAttention-3 from source and will be deprecated soon. Use _flash_3_hub or _flash_3_varlen_hub instead.

Try a backend temporarily

The attention_backend() context manager sets the process-wide active backend for the duration of the block and restores the previous backend when the block exits. Use it to try a backend for one call without leaving a permanent backend applied from set_attention_backend().

import torch
from diffusers import QwenImagePipeline, attention_backend

pipeline = QwenImagePipeline.from_pretrained(
    "Qwen/Qwen-Image", dtype=torch.bfloat16, device_map="cuda"
)
prompt = """
cinematic film still of a cat sipping a margarita in a pool in Palm Springs, California
highly detailed, high budget hollywood movie, cinemascope, moody, epic, gorgeous, film grain
"""

with attention_backend("_flash_3_hub"):
    image = pipeline(prompt).images[0]

Most attention backends work with torch.compile. Whether that speeds up your pipeline depends on the model and backend. See Precision and compilation.

Trusting remote kernels

Hub backends need the Kernels library first, and then Diffusers fetches the Hub kernel on first use.

Hub attention backends download compute kernels with the Kernels library and run them locally. Most Hub attention names (_flash_3_hub, flash_hub, and the other *_hub backends) resolve to the kernels-community organization. That organization is a trusted publisher in Kernels, so Diffusers loads those attention kernels without setting DIFFUSERS_TRUST_REMOTE_KERNELS.

The SageAttention Hub backends (sage_hub and sage_blackwell_hub) load from the SageAttention organization instead. Set DIFFUSERS_TRUST_REMOTE_KERNELS=true to use them.

Other kernel-backed features such as GGUF and Nunchaku Lite can pull kernels from publishers outside kernels-community. Those paths stay blocked unless you opt in with DIFFUSERS_TRUST_REMOTE_KERNELS. When set, Diffusers forwards trust_remote_code=True to Kernels so untrusted publishers can load too.

export DIFFUSERS_TRUST_REMOTE_KERNELS=true

Only enable this after inspecting the kernel repository. Without it, loading a kernel from an untrusted publisher raises an error. Diffusers performs this check itself, so it also applies to kernels<0.14.0, which predates the trust_remote_code argument. Setting DIFFUSERS_DISABLE_REMOTE_CODE=true disables remote code globally and takes precedence over DIFFUSERS_TRUST_REMOTE_KERNELS.

Checks

The attention dispatcher can run debugging checks before each dispatched attention call. Which checks run depends on the constraints registered for the active backend.

  1. Device checks verify that query, key, and value tensors live on the same device.
  2. Data type checks, where registered, confirm matching dtypes and often require bfloat16 or float16.
  3. Shape checks validate tensor dimensions and prevent mixing attention masks with causal flags.

Enable checks with the DIFFUSERS_ATTN_CHECKS environment variable. Checks add overhead, so they are disabled by default.

export DIFFUSERS_ATTN_CHECKS=yes

With checks on, Diffusers runs those constraints before every dispatched attention call. The low-level example below calls dispatch_attention_fn directly. Pipeline inference does not need that import. It only needs the backend set via set_attention_backend() or attention_backend().

import torch
from diffusers.models.attention_dispatch import attention_backend, dispatch_attention_fn

query = torch.randn(1, 10, 8, 64, dtype=torch.bfloat16, device="cuda")
key = torch.randn(1, 10, 8, 64, dtype=torch.bfloat16, device="cuda")
value = torch.randn(1, 10, 8, 64, dtype=torch.bfloat16, device="cuda")

try:
    with attention_backend("flash"):
        output = dispatch_attention_fn(query, key, value)
        print("✓ Flash Attention works with checks enabled")
except Exception as e:
    print(f"✗ Flash Attention failed: {e}")

Available backends

Refer to the table below for a complete list of available attention backends and their variants. Diffusers checks package availability and version pins when you enable a backend.

Backend Name Family Description Prerequisite
native PyTorch native Default backend using PyTorch's scaled_dot_product_attention None
flex FlexAttention PyTorch FlexAttention torch>=2.5.0
_native_cudnn PyTorch native CuDNN-optimized attention CUDA + CuDNN
_native_efficient PyTorch native Memory-efficient attention None beyond PyTorch
_native_flash PyTorch native PyTorch's FlashAttention CUDA
_native_math PyTorch native Math-based attention (fallback) None
_native_npu PyTorch native NPU-optimized attention torch_npu
_native_xla PyTorch native XLA-optimized attention torch_xla>=2.2
flash FlashAttention FlashAttention-2 flash-attn>=2.6.3
flash_hub FlashAttention FlashAttention-2 from Hub kernels kernels>=0.12
flash_varlen FlashAttention Variable length FlashAttention flash-attn>=2.6.3
flash_varlen_hub FlashAttention Variable length FlashAttention from Hub kernels kernels>=0.12
aiter_fa2_hub AI Tensor Engine for ROCm FlashAttention-2 for AMD ROCm from Hub kernels (bfloat16) kernels>=0.12, ROCm
flash_4_hub FlashAttention FlashAttention-4 from Hub kernels kernels>=0.12.3
_flash_3 FlashAttention FlashAttention-3 (local; deprecated soon) Build FA3 from source
_flash_varlen_3 FlashAttention Variable length FlashAttention-3 (local; deprecated soon) Build FA3 from source
_flash_3_hub FlashAttention FlashAttention-3 from Hub kernels kernels>=0.12
_flash_3_varlen_hub FlashAttention Variable length FlashAttention-3 from Hub kernels kernels>=0.12
sage SageAttention Quantized attention (INT8 QK) sageattention>=2.1.1
sage_hub SageAttention Quantized attention (INT8 QK) from Hub kernels kernels>=0.12, DIFFUSERS_TRUST_REMOTE_KERNELS=true
sage_blackwell_hub SageAttention SageAttention3 FP4 attention for SM120 Blackwell GPUs from Hub kernels kernels>=0.12, DIFFUSERS_TRUST_REMOTE_KERNELS=true
sage_varlen SageAttention Variable length SageAttention sageattention>=2.1.1
_sage_qk_int8_pv_fp8_cuda SageAttention INT8 QK + FP8 PV (CUDA) sageattention>=2.1.1
_sage_qk_int8_pv_fp8_cuda_sm90 SageAttention INT8 QK + FP8 PV (SM90) sageattention>=2.1.1; SM90
_sage_qk_int8_pv_fp16_cuda SageAttention INT8 QK + FP16 PV (CUDA) sageattention>=2.1.1
_sage_qk_int8_pv_fp16_triton SageAttention INT8 QK + FP16 PV (Triton) sageattention>=2.1.1
xformers xFormers Memory-efficient attention xformers>=0.0.29

Xet Storage Details

Size:
12.1 kB
·
Xet hash:
ca18f7a01e6a709d785a189cbbf17bc7164c1b6860f306ff383bdf45e0295679

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.