diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index b6eabfd584bb..4d303c9b61e4 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 @@ -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), +} + + 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: @@ -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)) frame_index, height_index, width_index = [], [], [] image_height_index, image_width_index = [], [] @@ -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( @@ -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__( @@ -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 ) image_ids[image_positions] = block_ids 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}" + )