From 95c69bfb778ac60b38d7e6e40c13523a47b4b547 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 24 Sep 2026 16:15:08 +0000 Subject: [PATCH 1/4] feat: qwen image 21 tp plan --- .../transformers/transformer_qwenimage21.py | 85 ++++++++++++++++--- 1 file changed, 74 insertions(+), 11 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index b6eabfd584bb..4d850a0dc2d3 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import functools import math from typing import Any @@ -23,7 +24,7 @@ from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import logging from ...utils.peft_utils import apply_lora_scale -from ...utils.torch_utils import maybe_allow_in_graph +from ...utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph from ..attention import AttentionMixin, AttentionModuleMixin from ..attention_dispatch import dispatch_attention_fn from ..cache_utils import CacheMixin @@ -133,6 +134,42 @@ def apply_rotary_emb_qwen( return x_out.type_as(x) +# Copied from diffusers.models.transformers.transformer_qwenimage.apply_rotary_emb_qwen_neuron +def apply_rotary_emb_qwen_neuron(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor: + """ + Apply rotary embeddings to `x` using real-valued cos/sin, for backends without a complex dtype. + + Numerically equivalent to `apply_rotary_emb_qwen(..., use_real=False)`, which multiplies `x` by a complex + exponential. Neuron has no complex tensor support, so the rotation angles are carried as reals and cos/sin are + taken here instead. + + Args: + x (`torch.Tensor`): Query or key tensor to rotate, shape `[B, S, H, D]`. + freqs (`torch.Tensor`): Rotation angles, shape `[S, D // 2]`. + + Returns: + `torch.Tensor`: `x` with rotary embeddings applied. + """ + # Adjacent feature pairs (2k, 2k+1) share angle k, so each angle is repeated twice along the last dim; unsqueeze + # the head axis so the freqs broadcast over heads (this is what keeps it tensor-parallel-agnostic). + cos = torch.cos(freqs).repeat_interleave(2, dim=-1).unsqueeze(1) # [S, 1, D] + sin = torch.sin(freqs).repeat_interleave(2, dim=-1).unsqueeze(1) # [S, 1, D] + x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] + x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) # [B, S, H, D] + return (x.float() * cos + x_rotated.float() * sin).to(x.dtype) + + +# RoPE application is backend-dependent: the default path multiplies by a complex exponential, which Neuron and TPU +# cannot represent. On those backends `QwenImage21Rope` hands out rotation angles instead of complex freqs, and +# `apply_rotary_emb_qwen_neuron` takes cos/sin on device. Callers select by `device.type` and fall back to the +# default for any backend not listed here. +_ROPE_ANGLE_DEVICES = ("neuron", "tpu") +ROPE_PER_DEVICE = { + "cuda": functools.partial(apply_rotary_emb_qwen, use_real=False), + **dict.fromkeys(_ROPE_ANGLE_DEVICES, apply_rotary_emb_qwen_neuron), +} + + class QwenImage21TemporalTimesteps(nn.Module): r"""Sinusoidal timestep embedding. `cos` occupies the first half of the channels and `sin` the second.""" @@ -337,16 +374,20 @@ def _qwenimage21_prepare_qkv( key = attn.to_k(hidden_states) value = attn.to_v(hidden_states) - query = query.unflatten(-1, (attn.heads, -1)) - key = key.unflatten(-1, (attn.heads, -1)) - value = value.unflatten(-1, (attn.heads, -1)) + # Split by `head_dim` rather than by `attn.heads`: under tensor parallelism each rank holds only its share of the + # heads, while `attn.heads` and `attn.inner_dim` keep their full values. + head_dim = attn.inner_dim // attn.heads + query = query.unflatten(-1, (-1, head_dim)) + key = key.unflatten(-1, (-1, head_dim)) + value = value.unflatten(-1, (-1, head_dim)) query = attn.norm_q(query).to(value.dtype) key = attn.norm_k(key).to(value.dtype) if rotary_emb is not None: - query = apply_rotary_emb_qwen(query, rotary_emb, use_real=False) - key = apply_rotary_emb_qwen(key, rotary_emb, use_real=False) + apply_rope = ROPE_PER_DEVICE.get(query.device.type, ROPE_PER_DEVICE["cuda"]) + query = apply_rope(query, rotary_emb) + key = apply_rope(key, rotary_emb) if layer_cache is not None: if kv_cache_mode == "extract" and cache_write_slice is not None: @@ -674,10 +715,19 @@ def rope_params(self, index: torch.Tensor, dim: int, theta: int = 10000) -> torc freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) return torch.polar(torch.ones_like(freqs), freqs) + @lru_cache_unless_export(maxsize=None) + def _get_device_freqs(self, device: torch.device) -> list[torch.Tensor]: + """Return the per-axis freqs on `device`: complex exponentials, or rotation angles where complex is missing.""" + if device.type in _ROPE_ANGLE_DEVICES: + # `torch.angle` runs on CPU while the freqs are still complex; wrapping into (-pi, pi] is harmless because + # only cos/sin of the angle are used. + return [torch.angle(freq).to(device) for freq in self.freqs] + return [freq.to(device) for freq in self.freqs] + def forward( self, img_shapes: list[tuple[int, int, int]], image_pad_mask: torch.Tensor, device: torch.device ) -> torch.Tensor: - self.freqs = [freq.to(device) for freq in self.freqs] + freqs = self._get_device_freqs(torch.device(device)) frame_index, height_index, width_index = [], [], [] image_height_index, image_width_index = [], [] @@ -707,7 +757,7 @@ def forward( height_index[image_pad_mask] = torch.tensor(image_height_index, dtype=torch.long, device=device) width_index[image_pad_mask] = torch.tensor(image_width_index, dtype=torch.long, device=device) - return torch.cat([self.freqs[0][frame_index], self.freqs[1][height_index], self.freqs[2][width_index]], dim=-1) + return torch.cat([freqs[0][frame_index], freqs[1][height_index], freqs[2][width_index]], dim=-1) class QwenImage21Transformer2DModel( @@ -758,6 +808,18 @@ class QwenImage21Transformer2DModel( _skip_layerwise_casting_patterns = ["pos_embed", "norm"] _repeated_blocks = ["QwenImage21TransformerBlock"] _skip_keys = ["kv_cache"] + # Tensor-parallel plan: every block's attention and SwiGLU projections are separate, bias-free Linears, so each + # entry is a plain "colwise"/"rowwise" pair and no packed sharding is needed. The shared `modulation`, the input + # projections (`img_in`, `txt_in`), `norm_out` and `proj_out` stay replicated (intentionally absent here). + _tp_plan = { + "transformer_blocks.*.attn.to_q": "colwise", + "transformer_blocks.*.attn.to_k": "colwise", + "transformer_blocks.*.attn.to_v": "colwise", + "transformer_blocks.*.attn.to_out.0": "rowwise", + "transformer_blocks.*.img_mlp.proj": "colwise", + "transformer_blocks.*.img_mlp.gate_layer": "colwise", + "transformer_blocks.*.img_mlp.out": "rowwise", + } @register_to_config def __init__( @@ -833,9 +895,10 @@ def build_token_metadata( ) image_ids = torch.full_like(image_pad_mask, -1, dtype=torch.long) - block_ids = torch.repeat_interleave( - torch.arange(len(block_lengths), device=image_pad_mask.device), - torch.tensor(block_lengths, device=image_pad_mask.device), + # Built from the Python block lengths rather than with a tensor-repeats `repeat_interleave`, whose + # data-dependent output size some compiled backends (e.g. Neuron) cannot lower. + block_ids = torch.tensor( + [block for block, length in enumerate(block_lengths) for _ in range(length)], device=image_pad_mask.device ) image_ids[image_positions] = block_ids From 481f25dc86b89dbf7c2554866856c5751cc51dc2 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Thu, 24 Sep 2026 16:18:09 +0000 Subject: [PATCH 2/4] test: vallidated on neuron --- .../test_models_transformer_qwenimage21.py | 43 ++++++++++++++++++- 1 file changed, 42 insertions(+), 1 deletion(-) diff --git a/tests/models/transformers/test_models_transformer_qwenimage21.py b/tests/models/transformers/test_models_transformer_qwenimage21.py index f6d0a0004d0c..f22a8a1addef 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -13,6 +13,10 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os +import subprocess +import sys + import pytest import torch from torch.nn.attention.flex_attention import create_mask @@ -21,7 +25,7 @@ from diffusers.models.transformers.transformer_qwenimage21 import build_qwenimage21_block_causal_mask from diffusers.utils.torch_utils import randn_tensor -from ...testing_utils import enable_full_determinism, torch_device +from ...testing_utils import enable_full_determinism, is_tensor_parallel, require_torch_neuron, torch_device from ..testing_utils import ( AttentionTesterMixin, BaseModelTesterConfig, @@ -29,6 +33,7 @@ ModelTesterMixin, SingleFileTesterMixin, TaylorSeerCacheTesterMixin, + TensorParallelTesterMixin, TrainingTesterMixin, ) @@ -380,3 +385,39 @@ def pretrained_model_kwargs(self): @property def torch_dtype(self): return torch.bfloat16 + + +class TestQwenImage21TransformerTensorParallel(QwenImage21TransformerTesterConfig, TensorParallelTesterMixin): + """Tensor Parallel inference tests for QwenImage 2.1 Transformer (CUDA/XPU multi-accelerator).""" + + +def make_neuron_tp_spec(): + """Model spec consumed by the generic Neuron TP worker (``_neuron_tp_worker.py``). + + Returns ``(model_class, init_dict, cpu_inputs)``. Reuses the shared tester config so the spec never drifts from the + rest of the QwenImage 2.1 tests. + """ + config = QwenImage21TransformerTesterConfig() + return QwenImage21Transformer2DModel, config.get_init_dict(), config.get_dummy_inputs(device="cpu") + + +@is_tensor_parallel +@require_torch_neuron +class TestQwenImage21TransformerTensorParallelNeuron: + """Tensor Parallel inference test for QwenImage 2.1 Transformer on AWS Neuron. + + Neuron TP runs through ``torchrun`` with the ``"neuron"`` distributed backend, so it cannot use the + ``torch.multiprocessing``/NCCL spawn path of ``TensorParallelTesterMixin``. This launches the generic worker with + the QwenImage 2.1 model spec (``make_neuron_tp_spec``); the worker asserts the sharded output matches a + single-device reference, and the test checks its exit code. + """ + + def test_tensor_parallel_neuron_inference(self): + worker = os.path.join(os.path.dirname(__file__), "_neuron_tp_worker.py") + spec = "tests.models.transformers.test_models_transformer_qwenimage21:make_neuron_tp_spec" + cmd = [sys.executable, "-m", "torch.distributed.run", "--nproc_per_node=2", worker, spec] + result = subprocess.run(cmd, capture_output=True, text=True) + assert result.returncode == 0, ( + f"Neuron tensor-parallel worker failed (exit {result.returncode}).\n" + f"--- stdout ---\n{result.stdout}\n--- stderr ---\n{result.stderr}" + ) From 46b415ab96bacd7f8b9d97cd989b9639eb462b43 Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 25 Sep 2026 12:43:33 +0000 Subject: [PATCH 3/4] fix: lru_cache_unless_export leak --- .../transformers/transformer_qwenimage21.py | 25 +++++++++++-------- 1 file changed, 15 insertions(+), 10 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index 4d850a0dc2d3..98884c116ac7 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -24,7 +24,7 @@ from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import logging from ...utils.peft_utils import apply_lora_scale -from ...utils.torch_utils import lru_cache_unless_export, maybe_allow_in_graph +from ...utils.torch_utils import maybe_allow_in_graph from ..attention import AttentionMixin, AttentionModuleMixin from ..attention_dispatch import dispatch_attention_fn from ..cache_utils import CacheMixin @@ -161,9 +161,9 @@ def apply_rotary_emb_qwen_neuron(x: torch.Tensor, freqs: torch.Tensor) -> torch. # RoPE application is backend-dependent: the default path multiplies by a complex exponential, which Neuron and TPU # cannot represent. On those backends `QwenImage21Rope` hands out rotation angles instead of complex freqs, and -# `apply_rotary_emb_qwen_neuron` takes cos/sin on device. Callers select by `device.type` and fall back to the -# default for any backend not listed here. -_ROPE_ANGLE_DEVICES = ("neuron", "tpu") +# `apply_rotary_emb_qwen_neuron` takes cos/sin on device. Callers select by `device.type` (PyTorch/XLA reports TPU +# tensors as `"xla"`) and fall back to the default for any backend not listed here. +_ROPE_ANGLE_DEVICES = ("neuron", "xla") ROPE_PER_DEVICE = { "cuda": functools.partial(apply_rotary_emb_qwen, use_real=False), **dict.fromkeys(_ROPE_ANGLE_DEVICES, apply_rotary_emb_qwen_neuron), @@ -710,19 +710,24 @@ def __init__(self, theta: int, axes_dim: list[int]): torch.cat([self.rope_params(pos_index, dim, theta), self.rope_params(neg_index, dim, theta)], dim=0) for dim in axes_dim ] + # Per-device copies of `freqs`, kept on the instance so they are freed with the model. A class-level + # `lru_cache` would key on `self` and keep every instance's device freqs alive for the life of the process. + self._device_freqs: dict[torch.device, list[torch.Tensor]] = {} def rope_params(self, index: torch.Tensor, dim: int, theta: int = 10000) -> torch.Tensor: freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) return torch.polar(torch.ones_like(freqs), freqs) - @lru_cache_unless_export(maxsize=None) def _get_device_freqs(self, device: torch.device) -> list[torch.Tensor]: """Return the per-axis freqs on `device`: complex exponentials, or rotation angles where complex is missing.""" - if device.type in _ROPE_ANGLE_DEVICES: - # `torch.angle` runs on CPU while the freqs are still complex; wrapping into (-pi, pi] is harmless because - # only cos/sin of the angle are used. - return [torch.angle(freq).to(device) for freq in self.freqs] - return [freq.to(device) for freq in self.freqs] + if device not in self._device_freqs: + if device.type in _ROPE_ANGLE_DEVICES: + # `torch.angle` runs on CPU while the freqs are still complex; wrapping into (-pi, pi] is harmless + # because only cos/sin of the angle are used. + self._device_freqs[device] = [torch.angle(freq).to(device) for freq in self.freqs] + else: + self._device_freqs[device] = [freq.to(device) for freq in self.freqs] + return self._device_freqs[device] def forward( self, img_shapes: list[tuple[int, int, int]], image_pad_mask: torch.Tensor, device: torch.device From f679fe59b590b7f4426e68e9777940867861eb2b Mon Sep 17 00:00:00 2001 From: JingyaHuang Date: Fri, 25 Sep 2026 14:09:43 +0000 Subject: [PATCH 4/4] fix: use complex RoPE on TPU for Qwen Image 2.1 TorchTPU reports TPU tensors as "tpu" and supports complex dtypes, so only Neuron needs the angle-based RoPE path. Co-Authored-By: Claude Opus 5.5 --- .../models/transformers/transformer_qwenimage21.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index 98884c116ac7..4d303c9b61e4 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -159,11 +159,11 @@ def apply_rotary_emb_qwen_neuron(x: torch.Tensor, freqs: torch.Tensor) -> torch. return (x.float() * cos + x_rotated.float() * sin).to(x.dtype) -# RoPE application is backend-dependent: the default path multiplies by a complex exponential, which Neuron and TPU -# cannot represent. On those backends `QwenImage21Rope` hands out rotation angles instead of complex freqs, and -# `apply_rotary_emb_qwen_neuron` takes cos/sin on device. Callers select by `device.type` (PyTorch/XLA reports TPU -# tensors as `"xla"`) and fall back to the default for any backend not listed here. -_ROPE_ANGLE_DEVICES = ("neuron", "xla") +# RoPE application is backend-dependent: the default path multiplies by a complex exponential, which Neuron cannot +# represent. On those backends `QwenImage21Rope` hands out rotation angles instead of complex freqs, and +# `apply_rotary_emb_qwen_neuron` takes cos/sin on device. Callers select by `device.type`. Every other backend (CUDA, +# CPU, TPU, …) uses the default complex path; only backends without complex dtypes are listed here. +_ROPE_ANGLE_DEVICES = ("neuron",) ROPE_PER_DEVICE = { "cuda": functools.partial(apply_rotary_emb_qwen, use_real=False), **dict.fromkeys(_ROPE_ANGLE_DEVICES, apply_rotary_emb_qwen_neuron),