Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
88 changes: 78 additions & 10 deletions src/diffusers/models/transformers/transformer_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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 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),

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I don't think we need this kind of dict munging. Let's just do: "neuron": 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."""

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -669,15 +710,29 @@ 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)

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 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
) -> torch.Tensor:
self.freqs = [freq.to(device) for freq in self.freqs]
freqs = self._get_device_freqs(torch.device(device))

@sayakpaul sayakpaul Sep 30, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Maybe we lru_cache this? Diff below:

diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py
--- a/src/diffusers/models/transformers/transformer_qwenimage21.py
+++ b/src/diffusers/models/transformers/transformer_qwenimage21.py
@@ -710,24 +710,19 @@
             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)
 
+    @functools.lru_cache(maxsize=128)
     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 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]
+        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


frame_index, height_index, width_index = [], [], []
image_height_index, image_width_index = [], []
Expand Down Expand Up @@ -707,7 +762,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(
Expand Down Expand Up @@ -758,6 +813,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__(
Expand Down Expand Up @@ -833,9 +900,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
)
Comment on lines -836 to 907

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Let's keep it explicitly conditioned on neuron then. @DN6 WDYT?

image_ids[image_positions] = block_ids

Expand Down
43 changes: 42 additions & 1 deletion tests/models/transformers/test_models_transformer_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -21,14 +25,15 @@
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,
MemoryTesterMixin,
ModelTesterMixin,
SingleFileTesterMixin,
TaylorSeerCacheTesterMixin,
TensorParallelTesterMixin,
TrainingTesterMixin,
)

Expand Down Expand Up @@ -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}"
)
Loading