Skip to content

[core] Tensor parallelism for Qwen-Image-2.1 - #14865

Open
JingyaHuang wants to merge 4 commits into
huggingface:mainfrom
JingyaHuang:add-qwen-image-21-tp-plan
Open

JingyaHuang wants to merge 4 commits into
huggingface:mainfrom
JingyaHuang:add-qwen-image-21-tp-plan

Conversation

@JingyaHuang

@JingyaHuang JingyaHuang commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

Adds tensor-parallel support to QwenImage21Transformer2DModel.

(Builds on #14544 sharded checkpoint loading, merge that one first.)

Changes in transformer_qwenimage21.py:

  • _tp_plan: attention to_q/to_k/to_v and SwiGLU proj/gate_layer are colwise; to_out.0 and img_mlp.out are rowwise. The shared modulation, input projections, norm_out and proj_out stay replicated.
  • Split query/key/value into heads by head_dim instead of attn.heads, so it still works when each rank only holds part of the heads (same fix as Qwen-Image).
  • RoPE per device, following Qwen-Image: CUDA keeps the complex path; Neuron which have no complex dtype, use real cos/sin.
  • build_token_metadata builds the block ids from a Python list instead of a tensor-repeats repeat_interleave, which Neuron can't compile. This one was needed for 2.1 to run on Neuron at all, TP or not.
import torch
from torch import distributed as dist
from diffusers import QwenImage21Pipeline, QwenImage21Transformer2DModel, TensorParallelConfig

dist.init_process_group(backend="nccl")
rank, world_size = dist.get_rank(), dist.get_world_size()
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)

transformer = QwenImage21Transformer2DModel.from_pretrained(
    "Qwen/Qwen-Image-2.1",
    subfolder="transformer",
    torch_dtype=torch.bfloat16,
    parallel_config=TensorParallelConfig(tp_degree=world_size),
)
pipe = QwenImage21Pipeline.from_pretrained("Qwen/Qwen-Image-2.1", transformer=transformer, torch_dtype=torch.bfloat16)
pipe.text_encoder.to(device)
pipe.vae.to(device)

image = pipe(
    "A capybara wearing a wizard hat, oil painting",
    generator=torch.Generator().manual_seed(0),  # same seed on every rank
).images[0]
if rank == 0:
    image.save("qwen21_tp.png")
dist.destroy_process_group()

Run with torchrun --nproc-per-node 8 qwen21_tp.py.

Validation

Mode Neuron TPU CUDA
Eager ✅ ✅ ✅
Compile

Before submitting

  • Did you use an AI agent (Claude Code, Codex, Cursor, etc.) to help with this PR? If so:
    • Did you read the Coding with AI agents guide?
    • Did you run the self-review skill on the diff?
    • Did you share the final self-review notes in the PR description or a comment?
  • Did you read the contributor guideline?
  • Did you read our philosophy doc? (important for complex PRs)
  • Was this discussed/approved via a GitHub issue or the forum? Please add a link to it if that's the case.
  • Did you make sure to update the documentation with your changes? Here are the
    documentation guidelines, and
    here are tips on formatting docstrings.
  • Did you write any new necessary tests?
  • Are you the author (or part of the team) of the model/pipeline (only applicable for model/pipeline related PRs)?

Who can review?

Anyone in the community is free to review the PR once the tests have passed. Feel free to tag
members/contributors who may be interested in your PR.

@github-actions github-actions Bot added documentation Improvements or additions to documentation lora models tests pipelines hooks size/L PR with diff > 200 LOC labels Sep 24, 2026
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

@JingyaHuang
JingyaHuang force-pushed the add-qwen-image-21-tp-plan branch from cc669f0 to 12e2b9c Compare September 25, 2026 14:16
@JingyaHuang
JingyaHuang marked this pull request as ready for review September 25, 2026 14:17

@sayakpaul sayakpaul left a comment

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.

Left some comments on the Qwen specific stuff.

_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.

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

Comment on lines -836 to 907
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
)

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?

JingyaHuang and others added 4 commits October 2, 2026 16:07
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 <noreply@anthropic.com>
@JingyaHuang
JingyaHuang force-pushed the add-qwen-image-21-tp-plan branch from fb74251 to f679fe5 Compare October 2, 2026 16:15
@github-actions github-actions Bot added size/M PR with diff < 200 LOC and removed documentation Improvements or additions to documentation lora pipelines hooks size/L PR with diff > 200 LOC labels Oct 2, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

models size/M PR with diff < 200 LOC tests

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants