Skip to content

[WS1][CUDA][Qwen3-Next] MoE route/combine contract with a batch-invariant grouped expert GEMM (RFC #428 stack 5/8) - #503

Open
fusheng-ji wants to merge 59 commits into
RL-Align:test-qwennextfrom
fusheng-ji:feat/ws1-qwen3-next-moe-route-combine
Open

fusheng-ji wants to merge 59 commits into
RL-Align:test-qwennextfrom
fusheng-ji:feat/ws1-qwen3-next-moe-route-combine

Conversation

@fusheng-ji

@fusheng-ji fusheng-ji commented Oct 9, 2026 •

Copy link
Copy Markdown

Repository refactor update

Merged RL-Align/RL-Kernel:test-qwennext at a5a92ecf6cbf29fdcfc862846e53db5fa244c35e as requested in RFC #428. This PR still targets test-qwennext.

Implementation changes now live in the canonical rl_engine/backends, rl_engine/reference, rl_engine/runtime, and rl_engine/validation directories. Where present, model assembly and vLLM training bridges use rl_engine/models/qwen3_next and rl_engine/integrations/engines/train/vllm, respectively. Tests use tests/models/qwen3_next; evidence and validation commands use tools/validation/models and tools/validation/operators. Existing upstream compatibility entry points and pinned operator identifiers are preserved. CI commands and path filters follow the new layout, including the test-qwennext target.

Validation after migration: 337 passed, 73 skipped in the relevant regression suites. CUDA tests used extensions built from the matching migrated native sources on B200. Syntax, undefined-name/duplicate-import checks, formatting, and git diff --check passed.

Stack: #466 (merged) → #467 → #468 → #469 → this PR, base test-qwennext; merges after #469. GitHub shows the whole stack's commits; this PR's own diff is the two commits on top of #469. RFC #428 row moe_route_combine_contract, CUDA.

Status

  • Ready for review. No known acceptance failure.
  • The routed MoE is batch-invariant end to end, forward and backward. Routes, output, dx and dW of a token are bitwise unchanged whatever else is in the batch (8 probe tokens alone vs. first/last in 16–1024-token batches). The TP4 gate on the real layer-0 weights passes on all four ranks.
  • The expert GEMMs reuse vLLM's fused_moe Triton kernel at the batch-invariant tile vLLM itself uses (64/64/32, SPLIT_K=1), passed explicitly. Its per-route output is bitwise equal to the per-expert pinned-GEMM loop it replaces, so this is a speed change with no change in bits. The loop it replaces took 66 ms at 1024 tokens in a local run; the forward is now 2.04 ms.
  • Known: still 2–4× slower than vLLM's non-strict MoE forward (2.04 vs 0.87 ms at 1024 tokens). About 0.5 ms at any token count is the FP32 IEEE router GEMM, which the contract requires. The backward is a per-expert loop: 169 ms forward + backward at 1024 tokens, against 51 ms for Megatron-core + TE (not batch-invariant) and 1284 ms for HF eager. It is the next thing to speed up.

What

The contract, per stage (full table in docs/operators/qwen3-next-moe-route-combine.md):

Stage Order and dtype Atomics
router BF16 x, W upcast; one pinned vLLM batch-invariant GEMM, IEEE FP32 none
selection argsort(logits, descending, stable)[:, :10]: score descending, ties to the lower expert id none
weights fixed-order softmax and renormalisation over a fixed pairwise tree, FP32 none
experts per (token, route) row: gate_up GEMM → SwiGLU in FP32 → BF16 → down GEMM; no routed weight in the kernel, no sum over routes none, one writer per row
combine e0·w0, then + e_k·w_k for k = 1..9 in route order, FP32; one BF16 cast none
backward per expert in ascending id; dW slices written once; dx summed per token in ascending expert order none
  • rl_engine/integrations/qwen3_next_forward.py: shared_linear, shared_router, stable_top10_routes, combine_routes, grouped_route_linear, shared_moe. Fails closed on CPU or ROCm tensors, on any vLLM other than 0.30.0, on VLLM_TRITON_USE_TD (a different instruction stream) and on TF32 Triton dots.
  • rl_engine/integrations/qwen3_next_tp_blocks.py: TP4MoE, with the HF load/export of all 512 experts (gate_up interleaved per rank), a replicated router and shared-expert gate, and MCore TP attributes for the gradient norm.
  • rl_engine/integrations/qwen3_next_tp.py: Megatron-style copy/reduce autograd boundaries through collective_for_group, shared with the later TP4 mixers.
  • rl_engine/testing/tensor_identity.py: raw-bit identity (signed zero distinct, NaN/Inf rejected, dtype must match).
  • Runners: scripts/qwen3_next_moe_prior_art.py + plot, scripts/qwen3_next_tp_moe_check.py (torchrun, 4 GPUs).

Prior art & reuse decision

Decision: reuse vLLM's batch-invariant GEMM (router, backward) and its fused_moe Triton kernel (expert GEMMs); implement the FP32 routing, tie rule, combine, backward and TP4 boundaries. Every library in the table was measured. No other implementation is batch-invariant and has a backward. The batch-invariant ones (vLLM BI=1, SGLang deterministic) are inference-only and route in BF16, which picks a different expert set than FP64 routing for 11 of 256 tokens. Megatron-core + TE, VIME's own training path, routes in FP32 like this PR but changes its routing, output, dx and dW with batch size.

Library Implementation Routes BI Output BI dx / dW BI Tokens routed unlike FP64 Notes
vLLM 0.30.0 fused_topk + fused_experts, VLLM_BATCH_INVARIANT=0 no (≥256) no (≥256) no backward 11 / 256 tuned configs and BF16 router GEMM vary with token count
vLLM 0.30.0 same, VLLM_BATCH_INVARIANT=1 (router through linear_batch_invariant) yes yes no backward 11 / 256 inference only; routed weight applied in the GEMM epilogue on BF16 output, then moe_sum
FlashInfer 0.6.18.post1 cutlass_fused_moe no (≥256) no (≥256) no backward 11 / 256 routes from torch.topk on BF16 logits
transformers 5.17.0 Qwen3NextExperts eager (Verl FSDP, Vime/Miles HF path) yes yes no at 1024 11 / 256 cuBLAS shape heuristics; index_add_ in BF16
Megatron-core 0.16 + TE 2.16 MoELayer, TE grouped GEMM + fused permute, configured as VIME runs Qwen3-Next (FP32 router, all-to-all) no (≥256) no (1024) no / no at 1024 0 / 256 VIME's native training path; accurate routing, not batch-invariant
SGLang 0.5.21 Triton fused_moe, default no (≥256) no (≥256) no backward 11 / 256 tuned configs and BF16 router GEMM vary with token count
SGLang 0.5.21 Triton fused_moe, deterministic inference yes yes no backward 11 / 256 inference only; deterministic tile 64/64/32 as in vLLM's BI mode
RL-Kernel (this PR) shared_moe / TP4MoE yes yes yes / yes 0 / 256

Results (B200, H = 2048, 512 experts, top-10, width 512, BF16)

Qwen3-Next routed MoE prior art

RL-Kernel Megatron + TE vLLM BI=1 SGLang det. vLLM BI=0 SGLang default FlashInfer HF eager
rel. L2 error vs FP64 3.9e-3 4.6e-3 7.0e-2 7.0e-2 7.0e-2 7.0e-2 7.0e-2 7.0e-2
forward, 1 / 64 / 1024 / 4096 tokens (ms) 1.39 / 1.69 / 2.04 / 3.05 3.8 / 9.7 / 11.7 / 12.1 0.39 / 0.66 / 0.87 / 1.05 0.35 / 0.66 / 0.84 / 1.37 0.38 / 0.66 / 0.87 / 1.06 0.34 / 0.61 / 0.83 / 1.35 0.37 / 0.72 / 0.90 / 1.12 2.1 / 62 / 89 / 87
forward + backward, 64 / 1024 tokens (ms) 114 / 169 45 / 51 – – – – – 877 / 1284
  • The accuracy gap is routing, not GEMM precision: with BF16 router logits, 11 of 256 tokens select a different expert set than FP64 routing; with an FP32 router (this PR, Megatron-core) none do.
  • TP4, real layer-0 weights: all 12 cases pass on 4 ranks: HF round trip of all 512 experts, identical routes on every rank, full == chunked and reversed-order batches at 8/64/256/1024 tokens, training forward == no-grad forward, identical replicated router and shared-gate gradients. TP4 block forward (routed + shared + all-reduce): 1.9–2.4 ms.

report.json, figure.png, tp4-moe/rank-*.json and a README with the exact commands are in docs/usage/evidence/qwen3-next-moe-route-b200/. The prior-art report was measured from a clean clone of c58816b in two Python environments (the training stack with Megatron-core/TE, and one with SGLang 0.5.21); the TP4 gate at 2e950f3, whose MoE code is identical. SGLang's compiled sgl_kernel cannot load next to torch 2.13, so its two kernels on this path are replaced by SGLang's own Triton moe_sum_reduce and JIT moe_align_block_size; any other sgl_kernel call fails. Latencies are CUDA-event medians with the candidate order reversed every iteration.

Tests

python setup.py build_ext --inplace
python -m pytest -q tests/test_tensor_identity.py tests/test_qwen3_next_forward_contract.py tests/test_qwen3_next_tp_blocks.py
TRITON_F32_DEFAULT=ieee python -m pytest -q tests/check_qwen3_next_forward.py        # needs vLLM 0.30.0
TRITON_F32_DEFAULT=ieee torchrun --nproc-per-node 4 scripts/qwen3_next_tp_moe_check.py \
    --checkpoint <Qwen3-Next-80B-A3B-Instruct> --output <dir>
TRITON_F32_DEFAULT=ieee python scripts/qwen3_next_moe_prior_art.py --out <dir>/report.json
Check Result
CPU: tensor identity, forward contract, TP blocks (+ test_framework_operator_integrations in its own process) 43 passed
check_qwen3_next_forward.py (GEMM batch/chunk/reorder, route ties, combine order, MoE VJP vs FP64, route and output BI, grouped == per-expert bitwise at widths 128 and 512, fail-closed TD path) 23 passed
TP4 gate, 4 × B200 12 / 12 cases, 4 ranks
prior-art runner report above

CI: the CPU tests join ci.yml. check_qwen3_next_forward.py imports vLLM, so it joins the self-hosted Qwen3-Next-provider-GPU workflow with TRITON_F32_DEFAULT=ieee.

Scope

CUDA, TP4 × EP1 (all 512 experts on every rank). ROCm, EP > 1 and other TP sizes are separate claims. Model-level parity uses this block but belongs to full_model_chain.

`test_registry_dispatches_rms_norm` asserts that the registry resolves
`rms_norm` to `RMSNormCudaOp` whenever CUDA and the compiled kernels are
both present, but `OpBackend` had no CUDA member for this operator and the
CUDA priority list contained only `PYTORCH_NATIVE_RMS_NORM`, so the assert
could never hold. The test therefore fails on any CUDA machine that builds
the native extension, and only passes when `_HAS_CUDA_RMSNORM` is false --
which is why an unbuilt CI has not caught it.

`RMSNormCudaOp` is already a first-class backend elsewhere: it is the
`"cuda"` candidate in `gtest/operator_specs.py` and is used directly by
`attention_preprocess.py`. Only the registry was missing it.

Add `OpBackend.CUDA_RMS_NORM` and put it ahead of the PyTorch reference in
the CUDA priority list. Because `_load_backend` only catches import errors
and this module imports cleanly without `_C`, a CUDA-first list would
otherwise hand out an op that raises at call time on an unbuilt install; so
`RMSNormCudaOp.__init__` now validates the extension and its three symbols,
matching `_require_cuda_activation` in the activation ops. The registry
already treats a backend whose construction raises as unavailable, so the
list degrades to `NativeRMSNormOp` as before.

Only the `cuda` priority map changes; rocm/musa/cpu/npu are untouched.

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130) with the
extension rebuilt from source:

  pytest tests/test_rms_norm.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

tests/test_rms_norm.py passes, including test_registry_dispatches_rms_norm, which
fails on the merge-base. The full suite gains no failure. The ignored file fails
to import on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…back

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ires

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The forward, backward-dx and both backward-dw launchers took the current
CUDA stream without switching to the input's device, so a tensor on cuda:1
while cuda:0 is current launched on the wrong GPU. Add a device guard on
the input's device in each launcher. Use at::cuda::OptionalCUDAGuard with
the headers included unconditionally, as activation.cu does: those
launchers compile in the ROCm build too, where the file's existing
c10::cuda::CUDAGuard stays inside the !USE_ROCM block. A two-device test
checks that the op runs on the input's device and matches the
single-device result.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Advances RFC RL-Align#428 C1 on the CUDA track. Claim level: L0 repeatable and L1
batch-invariant. L2 is NOT claimed -- see "On exactness against vLLM" below.

Qwen3-Next's decoder and final norms store a zero-centred weight and compute
`x * rstd * (1 + w)`, with the `1 +` applied in fp32 after the upcast. Folding
it into a bf16 weight beforehand rounds the offset away, so it has to reach the
kernel as a parameter rather than being pre-applied by the caller.

  NativeRMSNormOp            gains `weight_offset` (default 0.0)
  Qwen3NextRMSNormOp         subclasses it, overriding only weight_offset = 1.0
  Qwen3NextRMSNormGatedOp    GDN gated norm: plain w, weight multiply in fp32
  Qwen3NextRMSNormGatedHFOp  the transformers convention, kept as a witness

The offset is applied under `if cls.weight_offset:` rather than unconditionally,
because `0.0 + w` rewrites -0.0 to +0.0. torch.equal does not notice that, but a
bitwise comparison does, and the plain path must stay bit-for-bit what it was.
A test pins it at the bit level.

The gated pair exists because transformers and vLLM disagree on where the gated
norm's weight multiply happens, and the gap is not a ULP: on bf16 /
head_v_dim=128 they differ in 35% of elements with max|diff| = 6.25e-2.
Isolating the cast order alone reproduces the gap (5.3e-2), so the cast order
dominates rather than the reduction order. The two conventions share their
validation and normalization and differ only in a `_scale_by_weight` hook.

CUDA: `weight_offset` added to the forward and dx kernels, defaulting to 0.0 so
every existing caller and binding is unaffected. The dw kernel is untouched:
d/dw (offset + w) == d/dw w.

  weight_offset=1.0 vs an explicit fp32 (1 + w) weight   bitwise equal
  weight_offset=1.0 vs a bf16-folded (1 + w) weight      differs, as required
  default offset vs the previous kernel                  bitwise equal

The first line is the correctness argument: the in-kernel offset is the same
arithmetic as the fp32 reference, not an approximation. The second is a
regression guard -- if it ever passes, the offset has stopped being fp32.

On exactness against vLLM
-------------------------
Measured over 40 seeds (bf16, head_v_dim=128, 512 rows):

  ours vs forward_native            6/40 seeds differ, worst 1.56e-2
  ours vs forward_cuda             18/40 seeds differ, worst 3.91e-3
  forward_native vs forward_cuda   21/40 seeds differ, worst 1.56e-2

vLLM's own two paths are not bitwise equal to each other, so "bitwise equal to
vLLM" is undefined until a single provider is named. In fp32 the two paths
differ on ~36% of elements, every one by an fp32 ULP -- the tree shapes differ,
the semantics do not. What this reproduces is the convention (fp32 weight
multiply, single trailing cast); the residual is the reduction tree.

The reduction stays the repo's fixed 32-wide chunked sum, which is what buys
L1. It was introduced for NPU but is needed on CUDA too: over 20 seeds at
H=2048 in bf16 -- Qwen3-Next's own hidden_size and dtype -- a plain mean(-1)
broke slice invariance on 1 of 20 while the chunked reduction broke on 0 of 20.
Matching stock vLLM bitwise would mean adopting a reduction that is not itself
batch-invariant, i.e. trading L1 for L2.

tests/check_qwen3_next_norm_providers.py pins the dispatch facts and bounds the
gap, asserting magnitudes rather than equality so a vLLM bump that changes the
provider fails loudly. It is named `check_` rather than `test_`, following
tests/distributed/check_*.py: it imports real vLLM, and
tests/test_framework_operator_integrations.py asserts vllm is absent from
sys.modules, an invariant any collected test importing vLLM would break for the
whole session.

Also: `rl_engine/_C.pyi` updated for the new `weight_offset` argument (CI runs
mypy against it), `tests/test_qwen3_next_norm.py` added to the CI test list in
.github/workflows/ci.yml, and an operator page added per
docs/operators/README.md ("the documentation page is part of the operator
contract").

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130, triton
3.7.1, vllm 0.30.0, transformers 5.17.0):

  pytest tests/test_qwen3_next_norm.py -q
  pytest tests/check_qwen3_next_norm_providers.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

Both new files pass and the full suite gains no failure. The ignored file fails
to import on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… VJPs

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Adds the CUDA kernel behind `Qwen3NextRMSNormGatedOp`, and registers both it and
the zero-centred decoder norm from the previous commit.

  y = x * rstd * (weight_offset + weight) * act(gate)

Every multiply is fp32 with one cast at the store, matching vLLM's RMSNormGated
with norm_before_gate=True and group_size=None -- the only configuration the GDN
block constructs. Other configurations are rejected rather than approximated.
`expf` is used rather than the `__expf` intrinsic: the fast intrinsic trades
accuracy for speed and would move the result away from the fp32 reference.

Forward reuses the existing block_reduce_sum / choose_threads(H), so for a fixed
H the reduction tree is independent of the row count -- that is the L1 guarantee
-- and `rstd` comes out bitwise identical to the ungated kernel for the same x,
which is asserted.

Backward is assembled from deterministic pieces:
  dx      new kernel, the ungated dx with (w + offset) -> (w + offset) * act(z)
  dweight reuses rmsnorm_dweight_rows_fp32 + the ascending-row fp32 left fold
  dgate   row-local and reduction-free, fp32 in the wrapper

Both backwards route through one `_fold_dweight_rows` helper so the file keeps a
single left-fold entrypoint, which tests/test_vjp_fp32.py pins, and one
`_require_cuda_symbols` helper so the module has a single availability contract.

The gated op is deliberately NOT a subclass of RMSNormCudaOp: it takes an extra
required tensor, so it cannot stand in for one. Same reasoning as on the PyTorch
side, and stated in its docstring so the question is not reopened.

Registration: `rms_norm_gated` and `qwen3_next_rms_norm` in OP_SPECS (both with a
cuda-sm90 candidate, as every other reduction spec carries), operator_inputs
builders, OpBackend members and priority maps on all five platforms,
test_dispatch assertions, and the WS1 registered-ops set in
tests/test_ws1_gtest_gpu.py. Operator pages added per docs/operators/README.md,
which states the page is part of the operator contract.

Claim level: L0 repeatable and L1 batch-invariant. NOT L2 -- see
tests/check_qwen3_next_norm_providers.py, which pins the dispatch facts and
bounds the gap against vLLM rather than asserting equality.

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130, vllm 0.30.0):

  pytest tests/test_qwen3_next_norm.py -q
  pytest tests/check_qwen3_next_norm_providers.py -q
  python scripts/check_operator.py --op rms_norm_gated --candidate cuda \
      --device cuda --dtype bf16 --check-grad        -> pass_rate=1.0000
  python scripts/check_operator.py --op qwen3_next_rms_norm --candidate cuda \
      --device cuda --dtype bf16 --check-grad        -> pass_rate=1.0000
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

Both new files pass, both operators report pass_rate=1.0000 on the CUDA
candidate, and the full suite gains no failure. The ignored file fails to import
on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…t norms

Tests: the low-precision and CUDA-vs-golden comparisons hard-coded
atol=2e-2/rtol=1.6e-2, which is the contract's elementwise bf16 row. They now
resolve forward_accuracy for op_class="reduction" from tolerance_contract.json.
That loosens bf16 (5e-2/2e-2) and tightens fp16 (1e-3/1e-3, previously 2e-2);
the fp16 cases still pass. The vLLM provider-gap bounds in
tests/check_qwen3_next_norm_providers.py are labelled as gap bounds, not
contract thresholds.

Docs and docstrings: correct statements this branch had committed.
- The operator page no longer says the op is registered or prints a
  check_operator command; the gtest spec and registry entry arrive with the
  gated-norm PR.
- Withdrawn: "7 of 1048576 differ" (single seed, no script), the 35% / 5.3e-2
  cast-order isolation (an fp32 round-trip is a no-op), "needed on CUDA as
  well", and "matching vLLM means trading L1 for L2" (one unreproduced
  observation). Replaced with the scoped claim levels and the probe results.
- The provider check no longer claims to establish which path vLLM
  dispatches, or a per-element fp32 ULP bound.
- The module docstring is cut to the contract and links to the operator page.

Comments: why the kernel's offset add is guarded (signed zero, with the
pinning test), and why parameter_vjp_contributions_fp32 passes the offset.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…atistics

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ion fallback

The zero-centred norm was registered CUDA-first alongside the gated one, but only
rms_norm_gated had dispatch tests. Mirror them: pin the priority on all five
platforms, and assert that the inherited RMSNormCudaOp.__init__ symbol check makes
the registry fall through to the PyTorch reference.

Each assertion was checked to fail on CPU with the behaviour removed: PyTorch
first in the CUDA list, the ROCm entry dropped, and the __init__ check removed.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ctivation, offset

The bitwise rstd identity was asserted for one configuration only (bf16, H=128,
silu, offset 0). Parametrize it over fp32/fp16/bf16 x H in {128, 2048, 5120} x
{silu, sigmoid} x offset in {0.0, 1.0}, 36 cases, with the plain kernel taking
the same offset.

All 36 are expected to hold by construction: both kernels run the same
sum-of-squares loop, block_reduce_sum and choose_threads(H) launch shape, and
neither the gate nor the offset enters the statistic. The H values cover both
block sizes (128 and 512 threads) and two per-thread serial lengths. Not yet run
on a GPU; the cases skip without the compiled extension.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…16, 17)

R2-1 (private bf16 tolerances cited as the reduction row) was resolved in the
preceding merge, which switched the gated tests to c1's _forward_tol.

- R2-2 docs/operators/qwen3-next-rms-norm-gated.md: Accuracy now points to
  tolerance_contract.json (forward_accuracy, reduction x dtype); the 35% / 6.25e-2
  HF-vs-vLLM figure is labelled a one-off with its setup (one seed, B200, bf16,
  head_v_dim=128, 512 rows, eager); the 40-seed figures are attributed to
  tests/check_qwen3_next_norm_providers.py. Withdrawn, as on the c1 page: the
  cast-order isolation, the per-element fp32 ULP wording, and "fails closed (RFC
  section 6 item 7)" for configurations the op has no parameter for. Claim levels
  now read L1 (prefix slices; batch size, chunking, padding, permutation in the
  C3/C4 gates), L0 not separately tested; the rstd sweep is described.
- R2-7 csrc/ops.cpp: drop the dim and size TORCH_CHECKs that rmsnorm_check_weight
  already performs, in rmsnorm_forward, rmsnorm_backward_dx and
  rmsnorm_gated_check. Error messages for those cases now come from
  rmsnorm_check_weight; no test matches the removed messages.
- R2-8 workload_report() in rl_engine/testing/ws1_workload.py replaces the
  payload block duplicated in both gate scripts. full_model_evidence is read from
  the manifest: False for the Qwen3-Next norm manifest, None (not declared) for the
  Dense manifest, which previously reported a hard-coded False.
- R2-9 the gate scripts import validate_norm_dimensions at module level, and
  --manifest has a help string.
- R2-12 the four dispatch tests become two, parametrized over both norms. The
  fallback test now resolves for device="cuda", so the CUDA-first list is walked
  even on a CPU-only host, and asserts the CUDA backend was tried and failed.
  Before, on CPU, get_op went through the CPU list and the assertion held
  trivially.
- R2-13 SPDX headers on rl_engine/testing/qwen3_next_workload.py and
  tests/test_qwen3_next_workload.py.
- R2-14 ci/run_ws1_gtest.sh runs the four Qwen3-Next gates as explicit commands
  instead of a loop over interpolated script names.
- R2-16 Qwen3NextRMSNormGatedCudaOp takes activation as a constructor argument,
  validated before the extension check; the test no longer mutates the instance.
  "swish" stays as vLLM's alias for "silu", now asserted bitwise on CUDA.
- R2-17 the Python dgate path adds weight_offset only when it is nonzero, like the
  kernels, so a -0.0 weight keeps its sign in dgate.

New tests: activation rejected at construction (CPU), swish == silu (CUDA), dgate
signed zero (CUDA, 3 dtypes). Checked on CPU that the fallback and constructor
tests fail with the behaviour removed. The CUDA tests and the csrc change are not
yet run on a GPU.

Deferred: R2-6 (shared sum-of-squares device function), R2-10 (private validator
imports), R2-11.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… guard

The six rmsnorm launchers compiled on both CUDA and ROCm (plain forward and
dx, gated forward and dx, partial and reduce dweight) constructed a
non-optional c10::cuda::CUDAGuard. No other file the ROCm build compiles uses
that class outside a USE_ROCM guard, while activation.cu and
deterministic_attention.cu use at::cuda::OptionalCUDAGuard. Match the guard used
by activation.cu so the ROCm build relies only on patterns it already compiles.

The guard follows the same tensor as before (x, or partial_dw for the reduce
launcher). CUDA behaviour is unchanged: these tensors are always on a CUDA
device, so the optional guard always sets the device. The left-fold launcher
sits inside the USE_ROCM guard and keeps its c10::cuda::CUDAGuard. The ROCm
build is still untested.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…th the stacked C1 branch

Carried over from the former merge of the C1 branch into this one: the decoder-norm page again states the registration and the check_operator command, the gated tests use the contract's reduction tolerances through _forward_tol, and the reference docstrings name the gated page. Same final tree as before the re-stack.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The trainer-side reference for a Qwen3-Next rollout decode step: the recurrent
Gated DeltaNet update and the causal-conv1d state update that precedes it.
Claim level: L0 repeatable and L1 batch-invariant; L2 is not claimed.

Design notes, measurements and the deferral list are in
docs/design/ws1-c6-428-gdn-recurrent-replay.md.

Provider
--------
`fused_recurrent_gated_delta_rule_packed_decode`, not
`fused_sigmoid_gating_delta_rule_update`. A decode-only, non-speculative batch
returns early at qwen_gdn_linear_attn.py:1295-1307 into the packed path because
VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE defaults to true; the sigmoid-gating
kernel is the fallback for mixed prefill+decode and spec decode. The env
defaults that decide this are asserted, so a vLLM bump fails loudly.

Three places the kernel and the HF model disagree, transcribed from the kernel:

  * The gating is fused -- beta = sigmoid(b) and
    g = -exp(A_log) * softplus(a + dt_bias) are computed in fp32 inside the
    kernel with a threshold branch at 20, not as separate PyTorch ops.
  * No repeat_interleave: the kernel indexes i_h = i_hv // (HV // H).
  * The QK norm is an L2 norm over a plain sum, x / sqrt(sum(x*x) + 1e-6),
    dividing by sqrt rather than multiplying by rsqrt. scale hits q after the
    norm; k is never scaled.

A fourth candidate divergence turned out not to be one: prefill passes
use_qk_l2norm_in_kernel=False only because fused_post_conv_prep(apply_l2norm=True)
already normalized q/k. Both paths normalize exactly once.

Conv accumulation order is load-bearing: taps accumulate sequentially,
acc = acc + win[t] * w[t] from zero. With an fp32 cache that reproduces the
provider bitwise; a tree sum over the same four terms does not, and neither
does an FMA.

State ABI is mirrored: paged [num_blocks, HV, V, K] and
[num_blocks, dim, width-1], NULL_BLOCK_ID (<= 0) skipping, an fp32 accumulator,
and a store that rounds to the cache dtype -- fp32 or bf16, per
FUSED_GDN_STATE_DTYPES. Contractions run in the repo's fixed 32-wide chunk order
via one `_chunked_sum` primitive, which is what the L1 claim rests on.

Agreement, B200, Qwen3-Next dims (H=16, HV=32, K=V=128)
-------------------------------------------------------
  recurrent, fp32 state, B=1..64   max|d out| 1.5e-08..6.1e-05, state <= 3.0e-07
  recurrent, bf16 state, B=1..64   max|d out| 3.7e-09..3.1e-05, state <= 2.0e-03
  conv: the rolled state is BITWISE exact in every configuration
  conv, fp32 cache: output bitwise but for 15 of 524288 elements at B=64
  conv, bf16 cache: output agrees on ~63%, each disagreement one bf16 ULP
  L1: a sequence's output and state block are bitwise identical alone or at any
      position in a batch of 64, for both state dtypes

Where the provider rounds in the bf16-conv-cache case is not reproduced;
recorded as an open gap rather than guessed at.

Decode versus chunked prefill
-----------------------------
  T=8    max|d| 3.66e-04 fp32 state / 5.49e-04 bf16   (7.0e-03 / 1.0e-02 rel)
  T=64   4.88e-04 / 5.49e-04                          (6.6e-03 / 7.5e-03 rel)
  T=256  3.66e-04 / 3.66e-04                          (5.0e-03 / 5.0e-03 rel)

~0.5-1% relative and flat in T. The recurrence is contracting -- per-step decay
exp(g) averages ~0.47 -- so old rounding error is forgotten at roughly the rate
old signal is, and the fp32-vs-bf16 state drift saturates at ~2% relative on the
state (2.67e-03 at step 1, 1.52e-02 at 64, 2.47e-02 at 1024: a 16x longer run
past step 64 grows it 1.6x).

So the two paths are not bitwise, and ~1% relative on logits is still material
for RL importance ratios, but this is a bounded error rather than a divergence.
An fp32 recurrent state remains the right choice; the reason is the plateau, not
a blow-up. Separately, the chunked prefill kernel refuses fp32 q/k/v outright
(chunk.py:213), and causal_conv1d_update does not bounds-check
conv_state_indices under its default validate_data=False.

tests/check_gdn_recurrent_golden.py is named check_ rather than test_, following
tests/distributed/check_*.py: it imports real vLLM, and
tests/test_framework_operator_integrations.py asserts vllm is absent from
sys.modules.

  pytest tests/check_gdn_recurrent_golden.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

The new file passes and the full suite gains no failure.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Not covered, deliberately: speculative decode and MTP, a backward for the
recurrence, and the provider bridge. Reasons in the design doc.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…touched

test_null_block_id_is_skipped_by_both checked that both sides write zeros
for a NULL_BLOCK_ID row, but asserted "block untouched" only for the golden.
The provider returns before any store for state_idx <= 0
(fused_recurrent.py:299-303 in vllm 0.30.0); assert it on the provider's
state as well so the test name matches what it checks.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
ws1-gtest-gpu.yml: 158fc75 added four Qwen3-Next/GDN test files to the
pull_request and push path filters together with a step in
ci/run_ws1_gtest.sh that ran the two check_ files. 03176a3 moved that step
to qwen3-next-provider-gpu.yml but left the filters, so editing those files
triggered a RunPod job that never executed them. Remove the filters; the
file is back to its state on feat/cuda-qwen3-next-gated-rmsnorm.

qwen3-next-provider-gpu.yml: drop rl_engine/integrations/qwen3_next_gdn.py
and ci/run_qwen3_next_operator_gates.sh from the path filter. Neither file
exists on this branch.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The design note quoted provider-vs-golden figures (max|d out|, max|d state|,
conv mismatch counts) that no committed code printed, so they could not be
reproduced. scripts/ws1_gdn_provider_agreement.py prints them as JSON,
using tests/check_gdn_recurrent_golden.py's own input helpers and seeds
(imported, not copied):

- recurrent: per (batch, state dtype), max|d out|, max|d state| and
  bitwise mismatch counts, next to the bounds the check file asserts;
- conv: per (batch, cache dtype), output mismatch count, max|diff| and
  whether the rolled state is bitwise equal;
- the conv comparison again with triton.knobs.language.default_fp_fusion
  off, plus the fma.rn.f32 count in each compiled variant of the
  provider's conv-update kernel, to test whether the fp32-cache mismatches
  are FMA contraction.

Provenance (git commit and dirty flag, torch/triton/vllm versions, device)
is recorded with the results. Requires CUDA and vLLM; not run in CI.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Review after 03176a3 found nine statements in the note that the source
does not support. Fix them:

- §1: the decode path is chosen per engine step and by the model, not by
  env defaults alone. Tabulate decode-only / mixed / draft-token steps with
  their conv and recurrent kernels. Qwen3-Next's gqa_interleaved_layout=True
  makes the fused CUDA decode unsupported, so the layer falls back to triton
  and fused_gdn_decode_post_conv_mtp is unreachable (:1834). The env-default
  assertion guards the packed-decode default but not VLLM_GDN_DECODE_KERNEL.
- §2: the golden follows the kernel's arithmetic but is not a transcription
  everywhere: softplus uses log1p where the kernel uses log(1 + exp), and
  FLA_USE_FAST_OPS swaps in fast exp/log.
- §3: split NULL_BLOCK_ID semantics. The recurrent provider skips <= 0 and
  writes zeros; the conv provider skips only == null_block_id (0) and does
  not write the output row; both goldens skip <= 0 and write zeros.
- §4: replace the "measured" table, which had no runner, with the bounds the
  check file asserts and point to scripts/ws1_gdn_provider_agreement.py.
  State the fp32-cache conv mismatch and the FMA hypothesis the runner tests.
- §5: add the mixed decode-and-prefill caveat on the provider side, including
  vllm-project/vllm#49827, and state that validate_data=True does not check
  index values either.
- §6: restate the MTP deferral against RFC RL-Align#428 §2.2 and the actual
  Qwen3-Next MTP path.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
_git() returned the string "unavailable: ..." on failure, so
bool(_git("status", "--porcelain")) turned "git missing / not a checkout"
into git_dirty=True. Return None on failure and record git_dirty and
git_commit as null in that case, so the provenance distinguishes unknown
from dirty.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The §6 row deferred backward because "a naive BPTT through a sequential
recurrence is not batch-invariant". RFC RL-Align#428 does not ask backward to be:
§2.2 item 4 says backward need not match rollout, only be correct for the
replayed forward, and §9.1 makes the GDN backward its own work item, C7.
Cite those instead, matching the PR description.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
A run of the runner at d66e595 showed both of its FMA probes were ineffective:

- The fusion-off arm reused the fused kernel. Triton 3.7.1 keys its
  in-memory kernel cache on the launch kwargs; the default_fp_fusion knob
  is read in parse_options only after a cache miss. Flipping the knob
  in-process never missed, so all six recorded variants had
  enable_fp_fusion=True and the "off" counts were the "on" counts again.
  Run the fusion-off arm in a child process with TRITON_DEFAULT_FP_FUSION=0
  and a private TRITON_CACHE_DIR instead.
- fma.rn.f32 = 0 in the PTX does not rule out FMA. With fusion on, Triton
  emits plain mul.f32/add.f32 and lets ptxas contract them into FFMA; with
  fusion off it passes --fmad=false. Count plain and .rn PTX ops, and
  FFMA/FMUL/FADD in the SASS (Triton's bundled cuobjdump).

Report per-arm kernel variants under conv_kernels and add fusion_check,
which says whether the fusion-off variants were actually compiled without
fusion, so a silent no-op cannot pass for a result again.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The runner on d66e595 (B200, clean checkout) reproduced the recurrent
figures the note had withdrawn for lack of a runner, so restore them as a
per-batch table with their provenance and bitwise output-mismatch counts.

The conv counts quoted in test_conv_output_matches_provider_with_fp32_cache's
docstring (1 of 8192 at B=1, 15 of 524288 at B=64) predate the product-
rounding fix; the runner measures 0, 0, 1 and 5 for B = 1, 4, 17, 64.
Replace them, and record the bf16-cache counts in the note.

State plainly that the FMA explanation for the fp32-cache mismatches is not
determined: that run's fusion-off arm reused the fused kernel, and a PTX
without fma.rn.f32 cannot exclude ptxas contraction. The runner now tests
both (previous commit).

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…tches

The fixed runner (b54e50f) on 7622616, B200, settles the FMA hypothesis
left open in §4. Its fusion-off arm took effect: all six variants compiled
with enable_fp_fusion=False, and their PTX uses only .rn-qualified f32
mul/add, which ptxas may not contract. With contraction verifiably off,
the fp32-cache mismatch counts (0, 0, 1, 5 for B = 1, 4, 17, 64) and
max|diff| are identical to the fused run, and neither arm's SASS has an
FFMA. So FP contraction is ruled out; the actual cause stays open.

The note bases the conclusion on the A/B outcome, not on the opcode
counts: each variant's PTX has only two f32 multiplies, too few to be the
four tap products, so the counters do not show where the products are
computed.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
§5 said a mixed decode-and-prefill step leaves the decode row's bf16 output
equal. That holds only at TP1 head counts (H=16, HV=32). At TP4 per-rank
head counts (H=4, HV=8) the same step already moves one output element
(4.8e-05 relative) and 32,663 of 131,072 state elements (1.7e-07), with the
same numbers for all three prefill-bearing compositions.

Tabulate both head counts and state the measurement's limits: one
standalone layer from config.json, synthetic parameters and cache, single
process with no engine, scheduler or CUDA graph, metadata built by the test.
The 4-of-20 propagation figure is TP1-only and does not carry over to real
weights. Conv output and state matched bitwise, so the difference is in the
recurrent kernels. vllm-project/vllm#49827's two commits on 0.30.0 close the
decode-row gap at both head counts (scheduler part untested); with packed
decode disabled the decode row is bitwise equal in every composition.

Source: the GDN evidence branch's committed results (not yet published).
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Resolve the adjacent RMSNorm constructor and weight_offset insertion while preserving both the CUDA symbol guard and the zero-centred norm implementation.

Validation: 130 passed, 128 skipped across Qwen3-Next norm, RMSNorm, FP32 VJP and dispatch tests. The resolved tree is identical to the previous PR tip.
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Bring in PR RL-Align#467 at 81eae9e, including test-qwennext at 95914a8 after PR RL-Align#466 merged. The merged tree is identical to cf68fd9; this reconciles the stack ancestry without changing kernel or dispatch behavior.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…tions

Compares the zero-centred RMSNorm CUDA op with the PyTorch reference,
transformers, vLLM and FlashInfer (optional providers are skipped when
absent): error vs an FP64 golden, row invariance (256 rows alone vs a
4096-row batch, three seeds, bitwise) and forward/backward latency.
The plot script renders the report as one figure per op.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… branch

Brings in bec7e7a (RMSNorm API-version check for stale bindings) and the
evidence runner. The plain op now uses RL-Align#467's _require_cuda_rmsnorm; the
gated op keeps the generic _require_cuda_symbols, since the gated bindings
did not change signature.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…d RMSNorm

Adds the gated op (hidden 128, rows = tokens x heads) against the PyTorch
reference, transformers' cast-first Qwen3NextRMSNormGated and vLLM's
RMSNormGated, with an FP64 golden in vLLM's convention. Row inputs are now
generic, so the gate is sliced, differentiated and row-checked like x.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…6 GDN goldens

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Every sequence of a batch of 257 is compared, bitwise, alone and inside
sub-batches of 2, 7 and 64 at other offsets and a shuffled batch order,
with contiguous and shuffled cache blocks, fp32 and bf16 caches, with and
without bias and SiLU, over three seeds: output rows and rolled state
blocks are unchanged.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… on B200

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…rouped expert GEMM

RFC RL-Align#428 moe_route_combine_contract (CUDA).

- Routing: FP32 router through the pinned vLLM batch-invariant GEMM, a
  fixed-order softmax, stable descending argsort (ties go to the lower
  expert id) and FP32 renormalisation of the ten selected weights.
- Experts: per-route gate_up and down projections through vLLM 0.30.0's
  fused_moe Triton kernel with the batch-invariant tile passed explicitly
  (BLOCK 64/64/32, SPLIT_K=1); no routed-weight multiply and no sum inside
  the kernel, so every (token, route) row has one writer. Bitwise equal to
  the previous per-expert GEMM loop at TP4 and TP1 widths; fails closed if
  VLLM_TRITON_USE_TD is enabled.
- Combine: ten ordered FP32 multiply-adds and one BF16 cast; no atomics.
- Backward: per-expert pinned GEMMs into one dense gradient buffer; dx is
  summed in ascending expert order.
- TP4 block: HF load/export for all 512 experts, replicated router and
  shared-expert gate, Megatron-style copy/reduce autograd boundaries moved
  into qwen3_next_tp.py.
- Prior-art runner and plot comparing HF transformers, vLLM fused_moe
  (BI=0/1) and FlashInfer cutlass_fused_moe; TP4 real-weight gate script.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…nce (B200, 2e950f3)

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
@coderabbitai

coderabbitai Bot commented Oct 9, 2026 •

Copy link
Copy Markdown

Important

Review skipped

Auto reviews are disabled on base/target branches other than the default branch.

Please check the settings in the CodeRabbit UI or the .coderabbit.yaml file in this repository. To trigger a single review, invoke the @coderabbitai review command.

⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 4a7610fa-d65b-4e7a-b70c-08665710ed99

You can disable this status message by setting the reviews.review_status to false in the CodeRabbit configuration file.

Use the checkbox below for a quick retry:

  • 🔍 Trigger review
  • Autofix · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

…r-art runner

- Megatron-core 0.16 MoELayer with Transformer Engine grouped GEMM and fused
  permute, configured as VIME runs Qwen3-Next (FP32 softmax router, top-10,
  all-to-all dispatcher), including its backward.
- SGLang 0.5.21 Triton fused_moe in default and deterministic-inference mode.
  sgl_kernel cannot load next to torch 2.13, so its two kernels on this path
  are replaced by SGLang's own Triton moe_sum_reduce and JIT
  moe_align_block_size; any other sgl_kernel call fails loudly.
- --only and --merge combine reports from different Python environments.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ured (B200, c58816b)

Every library in the reuse table is now measured: Megatron-core 0.16 + TE 2.16 as VIME configures Qwen3-Next, and SGLang 0.5.21 in default and deterministic-inference mode. The TP4 gate evidence is unchanged (2e950f3, identical MoE code).

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…p as E402-exempt

The gate scripts insert the repo root into sys.path before importing
rl_engine and tools, which flake8 reports as E402. Mark those imports
with noqa, as the other tools/ scripts do.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Annotate the latency samples and accept the default-argument lambdas used as
timed calls. No behaviour change.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>

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