Repository navigation
Conversation
Make the two mixer forwards TP-correct so qwen3_5/3.6 can shard across ranks: - attention: per-rank q/kv head counts (div_even, KV replicates when num_kv_heads < tp_size); GQA-aware fused qkv built with explicit local output sizes (a uniform column split would halve a KV head); o_proj is row-parallel (all-reduce). - GDN: per-rank key/value heads matching LinearStatePool._linear_local_dims; conv/A_log/dt_bias at local sizes; out_proj row-parallel. TP=1 is a no-op (div_even(x,1)=x, no all-reduce), so single-GPU behavior is unchanged. The KV cache and attention backends are already TP-local.
…erts) Shard every weight to the local rank before the dense reader merges fused projections, so load_state_dict sees the TP-local tensors the layers build: - standard attention q/k/v/o, dense FFN gate/up/down, embed/lm_head via shard_tensor (handles GQA replication and the fp8 weight_scale_inv companions through its substring match); - GDN in_proj_qkv/conv1d head-segmented, in_proj_z/b/a + A_log/dt_bias by v-head, out_proj row-parallel; rules are role-agnostic so an fp8 128-block scale shards like its weight (per-head unit derived from the tensor); - Qwen3-VL vision tower (qkv/fc1 column, proj/fc2 row); - routed experts by intermediate: bf16 stacked [E,2I,H]/[E,H,I] and the block-fp8 resident stacks + their 128-block scales (tp must divide I/128). Drops the TP=1-only guards; TP=1 stays a no-op.
Allocate the fp8-block expert banks at the TP-local intermediate (matching the unquantized method and the sharded loader), and let the triton fp8-block MoE kernel run under TP>1 (its kernel-table guard rejected it). TP=1 unchanged.
The FlashML-org#601 bank system sizes each kernel's banks from layout(), which is already TP-local for the unquantized method and now for fp8-block too. Shard the full checkpoint pieces onto the local rank before pack(): gate/up families split the intermediate row axis, down and its block scale split the column axis, per-output-row globals replicate. No-op at tp=1, so every expert kind that reaches build_expert_banks under TP (its kernel must declare tp_ok) lands in the local banks instead of overflowing them.
NVFP4 linears carry a per-output-row fp16 global (weight_global [N]) and an optional per-tensor input_scale on top of the packed weight + block scale. The TP sharders handled the 2-D weight/scale but mis-handled those extras, so NVFP4 at tp>1 either mis-sharded the per-row global (column-parallel in_proj) or crashed indexing a 1-D tensor on dim 1 (row-parallel out_proj); the NVFP4 MoE bank layout also sized experts by the full intermediate, so each rank allocated full-size banks and TP never reduced memory. Rule applied throughout: a companion shards on the weight's sharded axis if it has that axis, else replicates. - qwen3_5_moe/weight.py: _module_leaf also strips .weight_global/.input_scale (companions resolve to the same GDN leaf); _shard_by_heads replicates a 0-d scalar or a companion whose shard dim doesn't exist (out_proj's [N] global); _shard_head_segments replicates a 0-d scalar (the 1-D [N] global head-shards with the weight via the existing shape[0] path). - moe/expert_banks.py: _shard_piece shards a per-row gate/up _global on dim1 only when it is per-row and divisible, else replicates (nvfp4's per-tensor scalar global). - models/loader.py: shard_tensor replicates a 0-d scalar and a 1-D per-row global on a row-parallel (un-sharded output) module. - quantization/moe/nvfp4.py: the Triton NVFP4 MoE bank layout/pack use cfg.local_intermediate so experts shard across ranks (the marlin/b12x CUDA-only kernels still size by full intermediate -- untested at tp>1).
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
This change enables tensor-parallel (TP > 1) serving of the Qwen3.5 and Qwen3.6 MoE models. Every class of weight is sharded across ranks at load time, so a checkpoint that does not fit on one GPU can be served across several. It covers bf16, fp8-block, and NVFP4 experts, and it is communication-backend agnostic (it works over either pynccl or torch-nccl).
What is sharded
Each weight is split along the dimension that matches how it is used in the forward pass:
q/k/vand the MLPgate/upprojections are column-parallel (split on their output dimension);o_projanddown_projare row-parallel (split on their input dimension); the embedding and thelm_headare split along the vocabulary dimension.Qwen3_5GatedDeltaNet): each rank keeps its local key/value heads;out_projis row-parallel (with an all-reduce); anddt_bias,A_log, andconv1dare allocated at the per-rank sizes. The linear-state pool is already TP-local, so it needs no change.Qwen3_5Attention): each rank keeps its local query and key/value heads (with GQA-aware key/value replication when there are fewer KV heads than ranks); theo_projis row-parallel.qkvandfc1are column-parallel (including their biases);projandfc2are row-parallel.build_expert_banks._shard_piece. Thegateanduptensors (and their scales) are split along their output rows; thedowntensor and its scale are split along their input columns. This path handles bf16 stacked experts, fp8-block experts (whose 128-blockgate_up_scaleanddown_scalecompanions are split the same way), and NVFP4 experts (packed e2m1 codes, a block scale, and the global scales).Implementation notes
weight_scale_inv(namedgate_up_scale/down_scalein the expert banks). An NVFP4 weight carries a per-output-rowweight_globalof shape[N], plus an optional per-tensorinput_scale. The sharding follows a single rule: a companion is split along the same axis as its weight when it has that axis, and is replicated otherwise. For example, a per-output-row global is split by head for a column-parallelin_proj; the same global is replicated for a row-parallelout_proj, becauseout_proj's output dimension is not sharded; and a per-tensor scalar is always replicated.layout()andpack()size the banks bycfg.local_intermediaterather than the full intermediate, so each rank allocates only its own slice. Without this, every rank would allocate full-size banks and tensor parallelism would not reduce expert memory. (The marlin and b12x NVFP4 kernels are CUDA-only and still size by the full intermediate; they are untested at TP > 1 and left as a follow-up.)LinearColParallelMergeddivides each output segment's width by the TP size internally, so the mixers pass the full segment widths to the layer and split the output using the local widths in the forward pass.Validation
--disable-pynccl(torch-nccl). This confirms the sharding is platform-agnostic.main, NVFP4 at TP=2 and fp8 at TP=4 were both re-checked and still serve coherently.Notes
DistributedCommunicator. Enabling RCCL/pynccl on AMD is a separate concern and is not part of this PR.moe_aligntorch fallback) from the ROCm-enablement PR. This PR is intended to sit on top of that stack.Test plan