From 0b94b5367081d81bac3513e2113b74d9547c3716 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Mon, 24 Aug 2026 06:01:42 +0000 Subject: [PATCH 1/6] propagate sage attention updates. --- src/diffusers/models/attention_dispatch.py | 2 +- tests/models/testing_utils/attention.py | 22 +++++++++++++++++++ tests/models/testing_utils/utils.py | 3 +++ .../test_models_transformer_qwenimage.py | 8 +++---- 4 files changed, 30 insertions(+), 5 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index e7cc20f580d4..d47ca9a44678 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -351,7 +351,7 @@ class _HubKernelConfig: AttentionBackendName.SAGE_HUB: _HubKernelConfig( repo_id="kernels-community/sage-attention", function_attr="sageattn", - version=1, + version=3, ), AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( repo_id="kernels-community/flash-attn4", diff --git a/tests/models/testing_utils/attention.py b/tests/models/testing_utils/attention.py index f31323d1bf52..6ff2390cb41e 100644 --- a/tests/models/testing_utils/attention.py +++ b/tests/models/testing_utils/attention.py @@ -94,6 +94,18 @@ ], ) +_PARAM_SAGE_HUB = pytest.param( + AttentionBackendName.SAGE_HUB, + id="sage_hub", + marks=[ + pytest.mark.skipif(not _CUDA_AVAILABLE, reason="CUDA is required for sage_hub backend."), + pytest.mark.skipif( + not is_kernels_available(), + reason="`kernels` package is required for sage_hub backend. Install with `pip install kernels`.", + ), + ], +) + # All backends under test. _ALL_BACKEND_PARAMS = [ _PARAM_NATIVE_CUDNN, @@ -101,12 +113,19 @@ _PARAM_FLASH_3_HUB, _PARAM_FLASH_VARLEN_HUB, _PARAM_FLASH_3_VARLEN_HUB, + _PARAM_SAGE_HUB, ] # Backends that perform non-deterministic operations and therefore cannot run when # torch.use_deterministic_algorithms(True) is active (e.g. after enable_full_determinism()). _NON_DETERMINISTIC_BACKENDS = {AttentionBackendName._NATIVE_CUDNN} +# Backends whose kernel cannot be traced into a single graph. Sage dispatches on the compute +# capability on every call (`torch.cuda.device_count()` returns a non-Tensor, which Dynamo +# rejects) and its arch-specific paths reach a Triton quantizer and torch ops that have no +# registered fake implementations. +_NO_FULLGRAPH_COMPILE_BACKENDS = {AttentionBackendName.SAGE_HUB} + def _skip_if_backend_requires_nondeterminism(backend): """Skip at runtime when torch.use_deterministic_algorithms(True) blocks the backend. @@ -419,6 +438,9 @@ def test_compile(self, backend, atol=1e-2, rtol=1e-2): if getattr(self.model_class, "_repeated_blocks", None) is None: pytest.skip("Skipping tests as regional compilation is not supported.") + if backend in _NO_FULLGRAPH_COMPILE_BACKENDS: + pytest.skip(f"Backend '{backend.value}' does not support fullgraph compilation.") + if backend == AttentionBackendName.NATIVE and not is_torch_version(">=", "2.9.0"): pytest.xfail( "test_compile with the native backend requires torch >= 2.9.0 for stable " diff --git a/tests/models/testing_utils/utils.py b/tests/models/testing_utils/utils.py index 07e4a38ddb21..9f2499ddca73 100644 --- a/tests/models/testing_utils/utils.py +++ b/tests/models/testing_utils/utils.py @@ -9,6 +9,9 @@ AttentionBackendName.FLASH_VARLEN_HUB, AttentionBackendName._FLASH_3_HUB, AttentionBackendName._FLASH_3_VARLEN_HUB, + # Sage attention quantizes QK to INT8 and PV to FP8/FP16, so it only accepts + # fp16/bf16 inputs and rejects the fp32 the test models default to. + AttentionBackendName.SAGE_HUB, } diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index 7a03a8fe2353..a301209bcf85 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage.py +++ b/tests/models/transformers/test_models_transformer_qwenimage.py @@ -259,7 +259,7 @@ class TestQwenImageTransformerAttention(QwenImageTransformerTesterConfig, Attent class TestQwenImageTransformerAttentionBackend(QwenImageTransformerTesterConfig, AttentionBackendTesterMixin): """Attention backend tests for QwenImage Transformer.""" - unsupported_attn_backends = ["flash_hub", "_flash_3_hub"] + unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub"] def get_dummy_inputs(self, batch_size: int = 2): inputs = super().get_dummy_inputs(batch_size=batch_size) @@ -289,9 +289,9 @@ class TestQwenImageTransformerContextParallelAttnBackends( ): """Context Parallel inference x attention backends tests for QwenImage Transformer""" - # QwenImage always passes a joint attention mask (text + image), which flash_hub and - # _flash_3_hub do not support. - unsupported_attn_backends = ["flash_hub", "_flash_3_hub"] + # QwenImage always passes a joint attention mask (text + image), which flash_hub, + # _flash_3_hub and sage_hub do not support. + unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub"] def get_dummy_inputs(self, batch_size: int = 1) -> dict[str, torch.Tensor]: inputs = super().get_dummy_inputs(batch_size=batch_size) From 691b43b53041b0456ec8b81be56673f2563338f2 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Mon, 31 Aug 2026 10:00:26 +0000 Subject: [PATCH 2/6] add sage blackwell. --- .ai/skills/diffusers-cli/run.md | 2 +- .../en/optimization/attention_backends.md | 1 + docs/source/en/using-diffusers/cli.md | 2 +- src/diffusers/models/attention_dispatch.py | 48 +++++++++++++++++++ tests/models/testing_utils/attention.py | 21 +++++++- tests/models/testing_utils/utils.py | 1 + .../test_models_transformer_flux.py | 8 ++++ .../test_models_transformer_qwenimage.py | 6 +-- 8 files changed, 83 insertions(+), 6 deletions(-) diff --git a/.ai/skills/diffusers-cli/run.md b/.ai/skills/diffusers-cli/run.md index 5eb56d826535..2f1d0f6cae54 100644 --- a/.ai/skills/diffusers-cli/run.md +++ b/.ai/skills/diffusers-cli/run.md @@ -108,7 +108,7 @@ Each entry calls `pipeline.load_lora_weights(, adapter_name=)`. A - `--cpu-offload {model, group}` — `model` uses `enable_model_cpu_offload`, `group` uses `enable_group_offload(offload_type="leaf_level", use_stream=True)`. Use `group` to fit a 9B+ model on a single A100. Onload target device comes from `--device-map` (must be a plain device string in this case). -- `--attention-backend {default, flash_hub, flash_varlen_hub, flash_4_hub, sage_hub}` — hub-hosted kernels, +- `--attention-backend {default, flash_hub, flash_varlen_hub, flash_4_hub, sage_hub, sage_blackwell_hub}` — hub-hosted kernels, auto-downloaded on first use. Failures (kernel not available, CUDA arch mismatch, network) raise a clear `SystemExit` listing the alternatives instead of silently reverting to the default. Only supported on transformer-based pipelines; UNet pipelines get a `logger.warning` and the flag is ignored. diff --git a/docs/source/en/optimization/attention_backends.md b/docs/source/en/optimization/attention_backends.md index 7602ddda1134..a33fbe89815f 100644 --- a/docs/source/en/optimization/attention_backends.md +++ b/docs/source/en/optimization/attention_backends.md @@ -164,6 +164,7 @@ Refer to the table below for a complete list of available attention backends and | `_flash_3_varlen_hub` | [FlashAttention](https://github.com/Dao-AILab/flash-attention) | Variable length FlashAttention-3 from kernels | | `sage` | [SageAttention](https://github.com/thu-ml/SageAttention) | Quantized attention (INT8 QK) | | `sage_hub` | [SageAttention](https://github.com/thu-ml/SageAttention) | Quantized attention (INT8 QK) from kernels | +| `sage_blackwell_hub` | [SageAttention](https://github.com/thu-ml/SageAttention) | SageAttention3 FP4 attention for SM120 Blackwell GPUs from kernels | | `sage_varlen` | [SageAttention](https://github.com/thu-ml/SageAttention) | Variable length SageAttention | | `_sage_qk_int8_pv_fp8_cuda` | [SageAttention](https://github.com/thu-ml/SageAttention) | INT8 QK + FP8 PV (CUDA) | | `_sage_qk_int8_pv_fp8_cuda_sm90` | [SageAttention](https://github.com/thu-ml/SageAttention) | INT8 QK + FP8 PV (SM90) | diff --git a/docs/source/en/using-diffusers/cli.md b/docs/source/en/using-diffusers/cli.md index 02be2bd2fd0e..8869e0bed032 100644 --- a/docs/source/en/using-diffusers/cli.md +++ b/docs/source/en/using-diffusers/cli.md @@ -134,7 +134,7 @@ Configure how the CLI loads model weights and custom pipeline code. `enable_auto_cpu_offload` as `memory_reserve_margin` (default `3GB`). Raise it when a large canvas runs out of memory mid-forward: the offloader keeps components resident while they fit, so on a high-VRAM card the default margin can leave too little room for the activations of a long video. -- `--attention-backend {default, flash_hub, flash_varlen_hub, flash_4_hub, sage_hub}` — Hub-hosted attention +- `--attention-backend {default, flash_hub, flash_varlen_hub, flash_4_hub, sage_hub, sage_blackwell_hub}` — Hub-hosted attention kernels, auto-downloaded on first use. Transformer-based pipelines only; ignored with a warning on legacy UNet pipelines. See [Attention backends](../optimization/attention_backends). - `--vae-tiling` / `--vae-slicing` — lower VAE decode VRAM. See diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index b873829bbd87..ee24bc075815 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -242,6 +242,7 @@ class AttentionBackendName(str, Enum): # `sageattention` SAGE = "sage" SAGE_HUB = "sage_hub" + SAGE_BLACKWELL_HUB = "sage_blackwell_hub" SAGE_VARLEN = "sage_varlen" _SAGE_QK_INT8_PV_FP8_CUDA = "_sage_qk_int8_pv_fp8_cuda" _SAGE_QK_INT8_PV_FP8_CUDA_SM90 = "_sage_qk_int8_pv_fp8_cuda_sm90" @@ -354,6 +355,11 @@ class _HubKernelConfig: function_attr="sageattn", version=3, ), + AttentionBackendName.SAGE_BLACKWELL_HUB: _HubKernelConfig( + repo_id="kernels-community/sage-blackwell", + function_attr="sageattn3_blackwell", + version=1, + ), AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( repo_id="kernels-community/flash-attn4", function_attr="flash_attn_func", @@ -473,6 +479,13 @@ def check_device_cuda(query: torch.Tensor, key: torch.Tensor, value: torch.Tenso return check_device_cuda +def _check_head_dim_64_or_128(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: + # The SM120 SageAttention3 kernel rejects head dims below 64 outright, fails to compile its + # Triton pre-pass on non-power-of-two dims, and silently falls back to SDPA at 256 and above. + if query.shape[-1] not in (64, 128): + raise ValueError(f"Query, key, and value must have a head dimension of 64 or 128, got {query.shape[-1]}.") + + def _check_qkv_dtype_match(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs) -> None: if query.dtype != key.dtype: raise ValueError("Query and key must have the same dtype.") @@ -535,6 +548,7 @@ def _check_attention_backend_requirements(backend: AttentionBackendName) -> None AttentionBackendName._FLASH_3_HUB, AttentionBackendName._FLASH_3_VARLEN_HUB, AttentionBackendName.SAGE_HUB, + AttentionBackendName.SAGE_BLACKWELL_HUB, AttentionBackendName.FLASH_4_HUB, AttentionBackendName.AITER_FA2_HUB, ]: @@ -4103,6 +4117,40 @@ def _sage_attention_hub( return (out, lse) if return_lse else out +@_AttentionBackendRegistry.register( + AttentionBackendName.SAGE_BLACKWELL_HUB, + constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_head_dim_64_or_128, _check_shape], +) +def _sage_attention_blackwell_hub( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attn_mask: torch.Tensor | None = None, + is_causal: bool = False, + scale: float | None = None, + return_lse: bool = False, + _parallel_config: "ParallelConfig" | None = None, +) -> torch.Tensor: + if attn_mask is not None: + raise ValueError("`attn_mask` is not supported for sage attention") + if return_lse: + # `sageattn3_blackwell` returns the output only, so there is no LSE to hand back. This + # also rules out context parallelism, hence `supports_context_parallel` is not set above. + raise ValueError("`return_lse` is not supported by the `sage_blackwell_hub` backend.") + if scale is not None and scale != query.shape[-1] ** -0.5: + # The kernel derives the softmax scale from the head dimension internally and silently + # swallows unknown kwargs, so a custom scale would be ignored rather than applied. + raise ValueError("A custom `scale` is not supported by the `sage_blackwell_hub` backend.") + + func = _HUB_KERNELS_REGISTRY[AttentionBackendName.SAGE_BLACKWELL_HUB].kernel_fn + # The kernel works on the HND layout, unlike the other Sage backends which take NHD. It also + # subtracts the per-token mean from `key` in place, so the transposed copies we build here + # double as protection for the caller's tensors. + query, key, value = (x.transpose(1, 2).contiguous() for x in (query, key, value)) + out = func(query, key, value, is_causal=is_causal) + return out.transpose(1, 2).contiguous() + + @_AttentionBackendRegistry.register( AttentionBackendName.SAGE_VARLEN, constraints=[_check_device_cuda, _check_qkv_dtype_bf16_or_fp16, _check_shape], diff --git a/tests/models/testing_utils/attention.py b/tests/models/testing_utils/attention.py index 6ff2390cb41e..1436a20b9bf9 100644 --- a/tests/models/testing_utils/attention.py +++ b/tests/models/testing_utils/attention.py @@ -36,6 +36,10 @@ # --------------------------------------------------------------------------- _CUDA_AVAILABLE = torch.cuda.is_available() +# Every build variant of `kernels-community/sage-blackwell` declares `archs: ["12.0a"]`, so the +# kernel only loads on SM120 (consumer/workstation Blackwell). `a` targets are architecture +# specific, so neither SM100 nor SM121 is covered. +_IS_SM120 = _CUDA_AVAILABLE and torch.cuda.get_device_capability() == (12, 0) _PARAM_NATIVE_CUDNN = pytest.param( AttentionBackendName._NATIVE_CUDNN, @@ -106,6 +110,20 @@ ], ) +_PARAM_SAGE_BLACKWELL_HUB = pytest.param( + AttentionBackendName.SAGE_BLACKWELL_HUB, + id="sage_blackwell_hub", + marks=[ + pytest.mark.skipif( + not _IS_SM120, reason="An SM120 Blackwell GPU is required for the sage_blackwell_hub backend." + ), + pytest.mark.skipif( + not is_kernels_available(), + reason="`kernels` package is required for sage_blackwell_hub backend. Install with `pip install kernels`.", + ), + ], +) + # All backends under test. _ALL_BACKEND_PARAMS = [ _PARAM_NATIVE_CUDNN, @@ -114,6 +132,7 @@ _PARAM_FLASH_VARLEN_HUB, _PARAM_FLASH_3_VARLEN_HUB, _PARAM_SAGE_HUB, + _PARAM_SAGE_BLACKWELL_HUB, ] # Backends that perform non-deterministic operations and therefore cannot run when @@ -124,7 +143,7 @@ # capability on every call (`torch.cuda.device_count()` returns a non-Tensor, which Dynamo # rejects) and its arch-specific paths reach a Triton quantizer and torch ops that have no # registered fake implementations. -_NO_FULLGRAPH_COMPILE_BACKENDS = {AttentionBackendName.SAGE_HUB} +_NO_FULLGRAPH_COMPILE_BACKENDS = {AttentionBackendName.SAGE_HUB, AttentionBackendName.SAGE_BLACKWELL_HUB} def _skip_if_backend_requires_nondeterminism(backend): diff --git a/tests/models/testing_utils/utils.py b/tests/models/testing_utils/utils.py index 9f2499ddca73..5070c443887d 100644 --- a/tests/models/testing_utils/utils.py +++ b/tests/models/testing_utils/utils.py @@ -12,6 +12,7 @@ # Sage attention quantizes QK to INT8 and PV to FP8/FP16, so it only accepts # fp16/bf16 inputs and rejects the fp32 the test models default to. AttentionBackendName.SAGE_HUB, + AttentionBackendName.SAGE_BLACKWELL_HUB, } diff --git a/tests/models/transformers/test_models_transformer_flux.py b/tests/models/transformers/test_models_transformer_flux.py index 53af9eedc50c..be76f892fc4c 100644 --- a/tests/models/transformers/test_models_transformer_flux.py +++ b/tests/models/transformers/test_models_transformer_flux.py @@ -253,6 +253,14 @@ class TestFluxTransformerAttention(FluxTransformerTesterConfig, AttentionTesterM class TestFluxTransformerAttentionBackend(FluxTransformerTesterConfig, AttentionBackendTesterMixin): """Attention backend tests for Flux Transformer.""" + def get_init_dict(self) -> dict[str, int | list[int]]: + # `sage_blackwell_hub` runs a kernel that only accepts head dims of 64 or 128, so widen the + # shared dummy config's `attention_head_dim` of 16. `axes_dims_rope` has to keep summing to it. + init_dict = super().get_init_dict() + init_dict["attention_head_dim"] = 64 + init_dict["axes_dims_rope"] = [16, 16, 32] + return init_dict + class TestFluxTransformerContextParallel(FluxTransformerTesterConfig, ContextParallelTesterMixin): """Context Parallel inference tests for Flux Transformer""" diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index a301209bcf85..5fcf37f6ff3f 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage.py +++ b/tests/models/transformers/test_models_transformer_qwenimage.py @@ -259,7 +259,7 @@ class TestQwenImageTransformerAttention(QwenImageTransformerTesterConfig, Attent class TestQwenImageTransformerAttentionBackend(QwenImageTransformerTesterConfig, AttentionBackendTesterMixin): """Attention backend tests for QwenImage Transformer.""" - unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub"] + unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub", "sage_blackwell_hub"] def get_dummy_inputs(self, batch_size: int = 2): inputs = super().get_dummy_inputs(batch_size=batch_size) @@ -290,8 +290,8 @@ class TestQwenImageTransformerContextParallelAttnBackends( """Context Parallel inference x attention backends tests for QwenImage Transformer""" # QwenImage always passes a joint attention mask (text + image), which flash_hub, - # _flash_3_hub and sage_hub do not support. - unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub"] + # _flash_3_hub and the sage hub backends do not support. + unsupported_attn_backends = ["flash_hub", "_flash_3_hub", "sage_hub", "sage_blackwell_hub"] def get_dummy_inputs(self, batch_size: int = 1) -> dict[str, torch.Tensor]: inputs = super().get_dummy_inputs(batch_size=batch_size) From 956f06191ea53d6a46b25aeb77baecb9b9adf6eb Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Mon, 21 Sep 2026 13:53:43 +0530 Subject: [PATCH 3/6] up --- src/diffusers/models/attention_dispatch.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 39537c8b0c04..8210b59e3724 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -47,7 +47,7 @@ is_xformers_available, is_xformers_version, ) -from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS +from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS, DIFFUSERS_TRUST_REMOTE_KERNELS from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph from ._modeling_parallel import gather_size_by_comm @@ -318,6 +318,7 @@ class _HubKernelConfig: wrapped_backward_attr: str | None = None wrapped_forward_fn: Callable | None = None wrapped_backward_fn: Callable | None = None + trust_remote_code: bool | list[str] = True # Registry for hub-based attention kernels @@ -351,14 +352,16 @@ class _HubKernelConfig: version=1, ), AttentionBackendName.SAGE_HUB: _HubKernelConfig( - repo_id="kernels-community/sage-attention", + repo_id="SageAttention/sage-attention", function_attr="sageattn", version=3, + trust_remote_code=["SageAttention/sage-attention"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False ), AttentionBackendName.SAGE_BLACKWELL_HUB: _HubKernelConfig( - repo_id="kernels-community/sage-blackwell", + repo_id="SageAttention/sage-blackwell", function_attr="sageattn3_blackwell", version=1, + trust_remote_code=["SageAttention/sage-blackwell"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False ), AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( repo_id="kernels-community/flash-attn4", @@ -739,11 +742,14 @@ def _maybe_download_kernel_for_backend(backend: AttentionBackendName) -> None: try: from kernels import get_kernel + trust_kwargs = {"trust_remote_code": config.trust_remote_code} if is_kernels_version(">=", "0.14.0") else {} + kernel_module = get_kernel( config.repo_id, revision=config.revision, version=config.version, user_agent={"diffusers": __version__}, + **trust_kwargs ) if needs_kernel: config.kernel_fn = _resolve_kernel_attr(kernel_module, config.function_attr) From c46fbef03ba8057b8ae194cd156e158d5bfc2cb6 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Mon, 21 Sep 2026 13:54:11 +0530 Subject: [PATCH 4/6] style. --- src/diffusers/models/attention_dispatch.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 8210b59e3724..547a2325f228 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -355,13 +355,13 @@ class _HubKernelConfig: repo_id="SageAttention/sage-attention", function_attr="sageattn", version=3, - trust_remote_code=["SageAttention/sage-attention"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False + trust_remote_code=["SageAttention/sage-attention"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False, ), AttentionBackendName.SAGE_BLACKWELL_HUB: _HubKernelConfig( repo_id="SageAttention/sage-blackwell", function_attr="sageattn3_blackwell", version=1, - trust_remote_code=["SageAttention/sage-blackwell"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False + trust_remote_code=["SageAttention/sage-blackwell"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False, ), AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( repo_id="kernels-community/flash-attn4", @@ -749,7 +749,7 @@ def _maybe_download_kernel_for_backend(backend: AttentionBackendName) -> None: revision=config.revision, version=config.version, user_agent={"diffusers": __version__}, - **trust_kwargs + **trust_kwargs, ) if needs_kernel: config.kernel_fn = _resolve_kernel_attr(kernel_module, config.function_attr) From edd2f8b87109b92be38c2f72da31ddc185c917e6 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 22 Sep 2026 08:21:14 +0530 Subject: [PATCH 5/6] address feedback --- src/diffusers/models/attention_dispatch.py | 21 ++++++++++++++++----- 1 file changed, 16 insertions(+), 5 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 547a2325f228..1d41cdd2fa75 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -50,6 +50,8 @@ from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS, DIFFUSERS_TRUST_REMOTE_KERNELS from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph from ._modeling_parallel import gather_size_by_comm +from huggingface_hub import get_organization_overview +from huggingface_hub.constants import HF_HUB_OFFLINE if TYPE_CHECKING: @@ -318,7 +320,6 @@ class _HubKernelConfig: wrapped_backward_attr: str | None = None wrapped_forward_fn: Callable | None = None wrapped_backward_fn: Callable | None = None - trust_remote_code: bool | list[str] = True # Registry for hub-based attention kernels @@ -355,13 +356,11 @@ class _HubKernelConfig: repo_id="SageAttention/sage-attention", function_attr="sageattn", version=3, - trust_remote_code=["SageAttention/sage-attention"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False, ), AttentionBackendName.SAGE_BLACKWELL_HUB: _HubKernelConfig( repo_id="SageAttention/sage-blackwell", function_attr="sageattn3_blackwell", version=1, - trust_remote_code=["SageAttention/sage-blackwell"] if DIFFUSERS_TRUST_REMOTE_KERNELS else False, ), AttentionBackendName.FLASH_4_HUB: _HubKernelConfig( repo_id="kernels-community/flash-attn4", @@ -742,10 +741,22 @@ def _maybe_download_kernel_for_backend(backend: AttentionBackendName) -> None: try: from kernels import get_kernel - trust_kwargs = {"trust_remote_code": config.trust_remote_code} if is_kernels_version(">=", "0.14.0") else {} + repo_id = config.repo_id + + if not HF_HUB_OFFLINE and not DIFFUSERS_TRUST_REMOTE_KERNELS: + publisher = repo_id.split("/")[0] + org_info = get_organization_overview(publisher) + if not getattr(org_info, "trustedKernelPublisher", False): + raise ValueError( + f"Backend '{backend.value}' loads `{config.repo_id}`, which is not published by a trusted kernel " + "publisher on the Hub, so loading it downloads and executes remote code. Set " + "`DIFFUSERS_TRUST_REMOTE_KERNELS=true` to allow it." + ) + + trust_kwargs = {"trust_remote_code": DIFFUSERS_TRUST_REMOTE_KERNELS} if is_kernels_version(">=", "0.14.0") else {} kernel_module = get_kernel( - config.repo_id, + repo_id, revision=config.revision, version=config.version, user_agent={"diffusers": __version__}, From ba24bdab1cd00431e02c0c04d1dfb78177f76512 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 22 Sep 2026 08:21:32 +0530 Subject: [PATCH 6/6] style. --- src/diffusers/models/attention_dispatch.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 1d41cdd2fa75..e907ddb5a01a 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -30,6 +30,9 @@ if torch.distributed.is_available(): import torch.distributed._functional_collectives as funcol +from huggingface_hub import get_organization_overview +from huggingface_hub.constants import HF_HUB_OFFLINE + from .. import __version__ from ..utils import ( get_logger, @@ -50,8 +53,6 @@ from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS, DIFFUSERS_TRUST_REMOTE_KERNELS from ..utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph from ._modeling_parallel import gather_size_by_comm -from huggingface_hub import get_organization_overview -from huggingface_hub.constants import HF_HUB_OFFLINE if TYPE_CHECKING: @@ -753,7 +754,9 @@ def _maybe_download_kernel_for_backend(backend: AttentionBackendName) -> None: "`DIFFUSERS_TRUST_REMOTE_KERNELS=true` to allow it." ) - trust_kwargs = {"trust_remote_code": DIFFUSERS_TRUST_REMOTE_KERNELS} if is_kernels_version(">=", "0.14.0") else {} + trust_kwargs = ( + {"trust_remote_code": DIFFUSERS_TRUST_REMOTE_KERNELS} if is_kernels_version(">=", "0.14.0") else {} + ) kernel_module = get_kernel( repo_id,