Repository navigation
[core] Tensor parallelism for Qwen-Image-2.1 - #14865
JingyaHuang wants to merge 9 commits into
Conversation
|
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. |
cc669f0 to
12e2b9c
Compare
sayakpaul
left a comment
There was a problem hiding this comment.
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), |
There was a problem hiding this comment.
I don't think we need this kind of dict munging. Let's just do: "neuron": apply_rotary_emb_qwen_neuron.
| 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 | ||
| ) |
There was a problem hiding this comment.
Let's keep it explicitly conditioned on neuron then. @DN6 WDYT?
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>
fb74251 to
f679fe5
Compare
sayakpaul
left a comment
There was a problem hiding this comment.
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):
(HF Jobs)
sayakpaul
left a comment
There was a problem hiding this comment.
Thanks! Just ran some tests and they're good. The changes in the DiT code don't lead to any changes in the outputs!
What does this PR do?
Adds tensor-parallel support to
QwenImage21Transformer2DModel, using the sharded checkpoint loadingfrom #14544.
Changes in `transformer_qwenimage21.p
_tp_plan: attentionto_q/to_k/to_vand SwiGLUproj/gate_layerare colwise;to_out.0and
img_mlp.outare rowwise. The shections,norm_outandproj_outstay replicated.head_diminsteal works when each rank holds only partof the heads (same fix as Qwen-Image).apply_rotary_emb_qwen_neuron, copied from Qwen-Image). All other devices keep the complex path. Theper-device freqs are now cached inste` on every forward.
block_idson Neuron: Neuron can't lower a tensor-repeatsrepeat_interleave, so on Neuron theblock ids are built from the Python bre unchanged. This was needed for 2.1to run on Neuron at all, TP or not.
Tests:
TensorParallelTesterMixinfor CUDA/XPU, and a Neuron TP=2 test that compares the sharded outputagainst a single-device reference.
Run with
torchrun --nproc-per-node 8 qwen21_tp.py.Validation
Before submitting
self-reviewskill on the diff?documentation guidelines, and
here are tips on formatting docstrings.
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.