Skip to content

feat(qwen3_5_moe): tensor-parallel weight sharding for multi-GPU serving - #637

Open
zhangnju wants to merge 5 commits into
FlashML-org:mainfrom
zhangnju:qwen3_5_moe_tp_sharding
Open

zhangnju wants to merge 5 commits into
FlashML-org:mainfrom
zhangnju:qwen3_5_moe_tp_sharding

Conversation

@zhangnju

@zhangnju zhangnju commented Oct 8, 2026

Copy link
Copy Markdown

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:

  • Dense linears: the attention q/k/v and the MLP gate/up projections are column-parallel (split on their output dimension); o_proj and down_proj are row-parallel (split on their input dimension); the embedding and the lm_head are split along the vocabulary dimension.
  • GDN mixer (Qwen3_5GatedDeltaNet): each rank keeps its local key/value heads; out_proj is row-parallel (with an all-reduce); and dt_bias, A_log, and conv1d are allocated at the per-rank sizes. The linear-state pool is already TP-local, so it needs no change.
  • Gated attention (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); the o_proj is row-parallel.
  • Vision tower (Qwen3-VL): qkv and fc1 are column-parallel (including their biases); proj and fc2 are row-parallel.
  • MoE experts: the unified expert banks are sharded along the intermediate dimension in build_expert_banks._shard_piece. The gate and up tensors (and their scales) are split along their output rows; the down tensor and its scale are split along their input columns. This path handles bf16 stacked experts, fp8-block experts (whose 128-block gate_up_scale and down_scale companions are split the same way), and NVFP4 experts (packed e2m1 codes, a block scale, and the global scales).

Implementation notes

  • Quantized formats carry companion tensors that must be sharded consistently with the weight they belong to. An fp8-block weight carries a 128-block weight_scale_inv (named gate_up_scale / down_scale in the expert banks). An NVFP4 weight carries a per-output-row weight_global of shape [N], plus an optional per-tensor input_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-parallel in_proj; the same global is replicated for a row-parallel out_proj, because out_proj's output dimension is not sharded; and a per-tensor scalar is always replicated.
  • The NVFP4 expert bank layout() and pack() size the banks by cfg.local_intermediate rather 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.)
  • LinearColParallelMerged divides 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

  • CUDA (2 x RTX 4090): Qwen3.5-9B at TP=2 produces greedy output that is byte-identical to TP=1 on two prompts, over both pynccl and --disable-pynccl (torch-nccl). This confirms the sharding is platform-agnostic.
  • ROCm (8 x R9700, gfx1201):
    • bf16 Qwen3.6-35B-A3B serves coherently at TP=4 and TP=8.
    • fp8 Qwen3.6-35B-A3B-FP8 produces byte-identical greedy output across TP=1 (offload), TP=2, and TP=4 on five prompts.
    • NVFP4 Qwen3.6-35B-A3B-NVFP4 serves coherently at TP=2. (This requires a TP-capable NVFP4 kernel at serve time; the stock triton NVFP4 kernel rejects TP > 1.)
  • After rebasing onto the current main, NVFP4 at TP=2 and fp8 at TP=4 were both re-checked and still serve coherently.

Notes

  • The change is communication-backend agnostic: all collectives go through DistributedCommunicator. Enabling RCCL/pynccl on AMD is a separate concern and is not part of this PR.
  • On RDNA, serving the MoE paths depends on the ROCm-7.14 / torch-2.14 compatibility fixes (notably the moe_align torch fallback) from the ROCm-enablement PR. This PR is intended to sit on top of that stack.

Test plan

  • CUDA: Qwen3.5-9B TP=2 output byte-identical to TP=1 (pynccl and torch-nccl)
  • ROCm: bf16 Qwen3.6-35B-A3B coherent at TP=4 and TP=8
  • ROCm: fp8 Qwen3.6-35B-A3B byte-identical across TP=1 / TP=2 / TP=4
  • ROCm: NVFP4 Qwen3.6-35B-A3B coherent at TP=2
  • marlin / b12x NVFP4 at TP > 1 (CUDA) — not covered by this PR

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

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant