Skip to content

feat(experimental): add bitwise Qwen3-MoE train–inference parity case based on GRPO + vLLM(TP>1) + FSDP - #430

Open
cboss6 wants to merge 2 commits into
Tencent-Hunyuan:mainfrom
cboss6:qwen3-moe-parity-tp4
Open

cboss6 wants to merge 2 commits into
Tencent-Hunyuan:mainfrom
cboss6:qwen3-moe-parity-tp4

Conversation

@cboss6

@cboss6 cboss6 commented Sep 9, 2026 •

Copy link
Copy Markdown
Contributor

Summary

This PR has been split into #504 and #505.

This PR introduces an experimental train–inference parity path for Qwen3-MoE.

In on-policy RL, numerical drift between vLLM rollout log-probabilities and differentiable actor replay can introduce artificial KL, move the importance ratio away from 1, and trigger clipping before any intended policy difference exists.

This experiment turns parity into a fail-closed runtime contract:

  • preserves sampled-token FP32 rollout log-probabilities;
  • replays the same prompt and response through differentiable teacher forcing;
  • requires strict torch.equal before backward;
  • synchronizes parity failures across FSDP ranks;
  • reports canonical token_count, torch_equal_fp32, mismatch_count, max_absdiff_fp32, k3_mean, and k3_max metrics;
  • aborts rather than training with mismatched log-probabilities.

Related PR

This PR can be merged before below PR merged:
#480

Supported case

The currently validated configuration is:

  • Model: Qwen3-30B-A3B
  • Actor: single-node FSDP4
  • Rollout: direct vLLM TP4
  • Compute: BF16
  • Log-probability comparison: FP32
  • Sampling: temperature=1, top-p=1, top-k=0
  • Prompt limit: 256 tokens
  • Response length: 1024 tokens
  • Batch size: 4 prompts
  • Samples per prompt: 8
  • Chunked prefill: enabled
  • Prefix caching: enabled
  • Data: shuffled GSM8K
  • Stack: Torch 2.11, Transformers 5.6, vLLM 0.22
  • Profile: public_reference

The normal Qwen3-MoE training path remains unchanged and does not load the experimental plugin.
Clipboard_Screenshot_1788974412

Post-update weight reload

The validation covers both the initial checkpoint and a non-zero optimizer update followed by FSDP → vLLM TP weight reload.

The dedicated synchronization path provides:

  • streamed full-weight staging;
  • strict TP payload cardinality;
  • per-tensor name, shape, dtype, numel, and SHA-256 validation;
  • consumed/loaded/missing/unexpected/duplicate-name receipts;
  • fail-stop RPC behavior;
  • transactional model-version commit;
  • prefix-cache reset only after all TP workers commit.

Validation

A clean two-rollout TP4 run passed both phases:

  • initial_checkpoint
  • post_update_reload

For both phases:

  • torch_equal_fp32 = true
  • mismatch_count = 0
  • max_absdiff_fp32 = 0.0
  • k3_mean = 0.0
  • k3_max = 0.0

The post-update reload synchronized 18,867 tensors across 114 buckets, committed the same model version on TP0–TP3, and verified that actor parameters changed.

Scope

The bitwise guarantee applies to selected-token FP32 forward log-probabilities on the validated hardware/software matrix.

It does not claim bitwise-identical gradients, optimizer states, or portability across arbitrary GPU architectures, CUDA versions, kernels, or vLLM releases.

Related Issue

Test Plan

Compatibility / Risk

Reviewer Notes

Checklist

  • I reviewed the changed code and removed unrelated/generated artifacts.
  • I updated tests, docs, and configs where needed, or explained why not.

Make the exact gate collective-safe, strengthen reload evidence, and keep custom Qwen3-MoE expert and o-projection layouts valid across vLLM sleep and native layerwise weight reload.

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

need review Ready and waiting for review

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant