Skip to content

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

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

JingyaHuang wants to merge 9 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, using the sharded checkpoint loading
from #14544.

Changes in `transformer_qwenimage21.p

  • _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 shections, norm_out and proj_out stay replicated.
  • Head split by head_dim insteal works when each rank holds only partof the heads (same fix as Qwen-Image).
  • RoPE on Neuron: Neuron has no coPE with real cos/sin(apply_rotary_emb_qwen_neuron, copied from Qwen-Image). All other devices keep the complex path. The
    per-device freqs are now cached inste` on every forward.
  • block_ids on Neuron: Neuron can't lower a tensor-repeats repeat_interleave, so on Neuron the
    block ids are built from the Python bre unchanged. This was needed for 2.1to run on Neuron at all, TP or not.

Tests: TensorParallelTesterMixin for CUDA/XPU, and a Neuron TP=2 test that compares the sharded output
against a single-device reference.

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(), d
device = torch.device(f"cuda:{rank}")
torch.cuda.set_device(device)

transformer = QwenImage21Transformer2
    "Qwen/Qwen-Image-2.1",
    subfolder="transformer",
    torch_dtype=torch.bfloat16,
    parallel_config=TensorParallelCon
)
pipe = QwenImage21Pipeline.from_pretrtransformer=transformer,torch_dtype=torch.bfloat16)
pipe.text_encoder.to(device)
pipe.vae.to(device)

image = pipe(
    "A capybara wearing a wizard hat,
    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.

Comment thread src/diffusers/models/transformers/transformer_qwenimage21.py Outdated
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
@JingyaHuang
JingyaHuang requested a review from sayakpaul October 7, 2026 23:49
Comment thread src/diffusers/models/transformers/transformer_qwenimage21.py Outdated

@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 one comment.

Ran the tests on HF Jobs using 2 A10g:
https://huggingface.co/jobs/sayakpaul/6ac71e14df2184ac91ac74c6

Also, ran the snippet and it ran fine (using 8 GPUs):

Image

(HF Jobs)

@sayakpaul
sayakpaul requested a review from DN6 October 8, 2026 10:02

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

Thanks! Just ran some tests and they're good. The changes in the DiT code don't lead to any changes in the outputs!

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