diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4634569cd..a73ecfa02 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,9 +5,9 @@ name: CI-Pipeline on: push: - branches: [ main, test ] + branches: [ main, test, test-qwennext ] pull_request: - branches: [ main, test ] + branches: [ main, test, test-qwennext ] permissions: contents: read @@ -75,7 +75,7 @@ jobs: run: | python -m pytest tests/runtime/test_dispatch.py -v PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest tests/ops/attention/test_attention_correctness.py -q -rs - python -m pytest tests/validation/operators/test_forward_invariance.py tests/contracts/test_tolerance_contract.py tests/validation/operators/test_ws1_workload.py tests/validation/operators/test_gradient_invariance.py tests/validation/operators/test_elementwise_inventory.py tests/validation/operators/test_four_judgment_matrix.py tests/validation/operators/test_op_checks.py tests/validation/operators/test_operator_inputs.py tests/benchmarks/test_profiler.py tests/validation/operators/test_kv_consistency.py tests/models/qwen3/test_ws1_qwen3_dense.py tests/validation/models/test_ws1_chain_integration.py -q + python -m pytest tests/validation/operators/test_forward_invariance.py tests/contracts/test_tolerance_contract.py tests/validation/operators/test_ws1_workload.py tests/validation/operators/test_gradient_invariance.py tests/validation/operators/test_elementwise_inventory.py tests/validation/operators/test_four_judgment_matrix.py tests/validation/operators/test_op_checks.py tests/validation/operators/test_operator_inputs.py tests/benchmarks/test_profiler.py tests/validation/operators/test_kv_consistency.py tests/models/qwen3/test_ws1_qwen3_dense.py tests/validation/models/test_ws1_chain_integration.py tests/models/qwen3_next/test_qwen3_next_norm.py -q - name: Run Cross-Configuration Contract Tests (CPU-safe) run: | diff --git a/.github/workflows/ws1-gtest-gpu.yml b/.github/workflows/ws1-gtest-gpu.yml index a0d8ac0c4..b853b9df0 100644 --- a/.github/workflows/ws1-gtest-gpu.yml +++ b/.github/workflows/ws1-gtest-gpu.yml @@ -10,7 +10,7 @@ name: WS1-gtest-GPU on: pull_request: - branches: [ main ] + branches: [ main, test-qwennext ] paths: - "rl_engine/validation/**" - "rl_engine/backends/**" @@ -28,6 +28,7 @@ on: - "tools/validation/operators/ws1_candidate_evidence.py" - "tests/validation/**" - "tests/models/**" + - "tests/ops/norm/**" - "tests/validation/operators/test_forward_invariance.py" - "tests/validation/operators/test_gradient_invariance.py" - "tests/validation/operators/test_four_judgment_matrix.py" @@ -36,7 +37,7 @@ on: - "ci/scripts/run_gpu_ci.sh" - ".github/workflows/ws1-gtest-gpu.yml" push: - branches: [ main ] + branches: [ main, test-qwennext ] paths: - "rl_engine/validation/**" - "rl_engine/backends/**" @@ -54,6 +55,7 @@ on: - "tools/validation/operators/ws1_candidate_evidence.py" - "tests/validation/**" - "tests/models/**" + - "tests/ops/norm/**" - "tests/validation/operators/test_forward_invariance.py" - "tests/validation/operators/test_gradient_invariance.py" - "tests/validation/operators/test_four_judgment_matrix.py" diff --git a/ci/scripts/run_ws1_gtest.sh b/ci/scripts/run_ws1_gtest.sh index cc47dd678..52caa1c7d 100755 --- a/ci/scripts/run_ws1_gtest.sh +++ b/ci/scripts/run_ws1_gtest.sh @@ -19,6 +19,7 @@ echo "[ws1-gtest] interpreter=$PY out=$OUT" "$PY" -m pytest -q \ tests/validation/operators/test_ws1_gtest_gpu.py \ + tests/models/qwen3_next/test_qwen3_next_norm.py \ tests/ops/attention/test_triton_batch_invariant_attention.py \ tests/validation/operators/test_four_judgment_matrix.py \ tests/validation/operators/test_ws1_candidate_evidence.py \ diff --git a/csrc/bindings/ops.cpp b/csrc/bindings/ops.cpp index 0c7c3e7f4..6fcf90752 100644 --- a/csrc/bindings/ops.cpp +++ b/csrc/bindings/ops.cpp @@ -263,14 +263,16 @@ void rmsnorm_forward_cuda( torch::Tensor weight, torch::Tensor y, torch::Tensor rstd, - double eps); + double eps, + double weight_offset); void rmsnorm_backward_dx_cuda( torch::Tensor dy, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, - torch::Tensor dx); + torch::Tensor dx, + double weight_offset); void rmsnorm_backward_partial_dw_cuda( torch::Tensor dy, @@ -299,7 +301,8 @@ static void rmsnorm_check_input(const torch::Tensor& x, const char* name) { std::vector rmsnorm_forward( torch::Tensor x, torch::Tensor weight, - double eps) + double eps, + double weight_offset) { rmsnorm_check_input(x, "x"); rmsnorm_check_input(weight, "weight"); @@ -312,7 +315,7 @@ std::vector rmsnorm_forward( auto y = torch::empty_like(x); auto rstd = torch::empty({T}, x.options().dtype(torch::kFloat32)); - rmsnorm_forward_cuda(x, weight, y, rstd, eps); + rmsnorm_forward_cuda(x, weight, y, rstd, eps, weight_offset); return {y, rstd}; } @@ -321,7 +324,8 @@ torch::Tensor rmsnorm_backward_dx( torch::Tensor dy, torch::Tensor x, torch::Tensor weight, - torch::Tensor rstd) + torch::Tensor rstd, + double weight_offset) { rmsnorm_check_input(dy, "dy"); rmsnorm_check_input(x, "x"); @@ -336,7 +340,7 @@ torch::Tensor rmsnorm_backward_dx( auto dx = torch::empty_like(x); - rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx); + rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx, weight_offset); return dx; } @@ -718,8 +722,14 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &det_gemm_db_transposed, "Batch-invariant deterministic GEMM backward in canonical [N,K] layout"); // registry RMSNorm - m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA"); - m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA"); + // API version 2 includes weight_offset in both entry points. + m.attr("rmsnorm_api_version") = 2; + m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA", + py::arg("x"), py::arg("weight"), py::arg("eps"), + py::arg("weight_offset") = 0.0); + m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA", + py::arg("dy"), py::arg("x"), py::arg("weight"), py::arg("rstd"), + py::arg("weight_offset") = 0.0); m.def("rmsnorm_backward_dw", &rmsnorm_backward_dw, "Deterministic RMSNorm backward dweight CUDA"); #if !defined(USE_ROCM) m.def( diff --git a/csrc/cuda/norm/rmsnorm.cu b/csrc/cuda/norm/rmsnorm.cu index b32bc5af4..595131522 100644 --- a/csrc/cuda/norm/rmsnorm.cu +++ b/csrc/cuda/norm/rmsnorm.cu @@ -108,7 +108,8 @@ __global__ void rmsnorm_fwd_kernel( float* __restrict__ rstd, int T, int H, - float eps + float eps, + float weight_offset ) { int row = blockIdx.x; int tid = threadIdx.x; @@ -135,10 +136,17 @@ __global__ void rmsnorm_fwd_kernel( __syncthreads(); - // Write y = x * rstd * weight. + // Write y = x * rstd * (weight_offset + weight). The offset is added in + // fp32 after the upcast: a zero-centred weight (Qwen3-Next, Gemma) must not + // have its "+1" folded into the low-precision weight beforehand, which would + // round the offset and break the bitwise contract. for (int col = tid; col < H; col += blockDim.x) { float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + // Guarded: `-0.0f + 0.0f` is +0.0f, so an unconditional add would flip the + // sign bit of -0.0 weights on the plain path. Pinned by + // tests/models/qwen3_next/test_qwen3_next_norm.py::test_cuda_plain_signed_zero_is_preserved. + if (weight_offset != 0.0f) wv += weight_offset; float out = xv * row_rstd * wv; store_from_float(y_row + col, out); } @@ -153,7 +161,8 @@ __global__ void rmsnorm_bwd_dx_kernel( const float* __restrict__ rstd, scalar_t* __restrict__ dx, int T, - int H + int H, + float weight_offset ) { int row = blockIdx.x; int tid = threadIdx.x; @@ -168,6 +177,7 @@ __global__ void rmsnorm_bwd_dx_kernel( float dyv = load_as_float(dy_row + col); float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; local_dot += dyv * wv * xv; } @@ -180,6 +190,7 @@ __global__ void rmsnorm_bwd_dx_kernel( float dyv = load_as_float(dy_row + col); float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; float out = r * dyv * wv - xv * coeff; store_from_float(dx_row + col, out); @@ -247,7 +258,8 @@ void rmsnorm_forward_cuda( torch::Tensor weight, torch::Tensor y, torch::Tensor rstd, - double eps + double eps, + double weight_offset ) { // Launch on x's device: the current CUDA stream belongs to the current // device, which need not be x's. @@ -270,7 +282,8 @@ void rmsnorm_forward_cuda( rstd.data_ptr(), T, H, - static_cast(eps) + static_cast(eps), + static_cast(weight_offset) ); }); }); @@ -282,7 +295,8 @@ void rmsnorm_backward_dx_cuda( torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, - torch::Tensor dx + torch::Tensor dx, + double weight_offset ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); @@ -303,7 +317,8 @@ void rmsnorm_backward_dx_cuda( rstd.data_ptr(), dx.data_ptr(), T, - H + H, + static_cast(weight_offset) ); }); }); diff --git a/docs/.nav.yml b/docs/.nav.yml index 29b5ead37..ebb5476de 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -29,6 +29,7 @@ nav: - operators/sampling.md - operators/det-gemm.md - operators/embedding.md + - operators/qwen3-next-rms-norm.md - Developer Guide: - contributing/README.md - Contributor Guide: contributing/contributor-guide.md diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..0d955ed18 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -32,4 +32,5 @@ Every operator page should include: - [Matmul](matmul.md) - [Sampling](sampling.md) - [Token Embedding](embedding.md) +- [Qwen3-Next RMSNorm (zero-centred)](qwen3-next-rms-norm.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/qwen3-next-rms-norm.md b/docs/operators/qwen3-next-rms-norm.md new file mode 100644 index 000000000..0e42a0c8b --- /dev/null +++ b/docs/operators/qwen3-next-rms-norm.md @@ -0,0 +1,183 @@ +# Qwen3-Next RMSNorm (zero-centred) + +## Summary + +Qwen3-Next's decoder and final norms store a **zero-centred** weight and compute +`x * rstd * (1 + w)` rather than `x * rstd * w`. The `1 +` is applied in fp32, +after the upcast — folding it into a low-precision weight beforehand rounds the +offset away and silently breaks any bitwise claim. + +This operator exists for RFC #428 C1 (Embedding / RMSNorm / residual / final norm +exactness) on the Qwen3-Next rollout-vs-replay path. The Gated DeltaNet block uses +a *different* weight convention and a different cast order; it is a separate +operator. + +Upstream references: `transformers` `Qwen3NextRMSNorm`, and vLLM's `GemmaRMSNorm`, +which `vllm/model_executor/models/qwen3_next.py` aliases as `Qwen3NextRMSNorm`. + +## Entry Point + +```python +from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormOp +from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp, rmsnorm_cuda + +y = Qwen3NextRMSNormOp().forward(x, weight, eps=1e-6) # reference +y = Qwen3NextRMSNormCudaOp().forward(x, weight, eps=1e-6) # CUDA + +# The offset is a kernel parameter, not a pre-pass: +y = rmsnorm_cuda(x, weight, eps=1e-6, weight_offset=1.0) +``` + +## Backends + +| Backend | Wrapper | Native symbol | Status | +| --- | --- | --- | --- | +| CUDA | `Qwen3NextRMSNormCudaOp` | `rl_engine._C.rmsnorm_forward` (`weight_offset=1.0`) | Supported | +| ROCm | — | — | Not implemented | +| PyTorch fallback | `Qwen3NextRMSNormOp` | — | Supported (WS1 gold) | + +`Qwen3NextRMSNormOp` subclasses `NativeRMSNormOp` and overrides only +`weight_offset`; the base applies it in fp32. + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `x` | `[..., H]` | fp32 / bf16 / fp16 | CUDA wrapper copies non-contiguous inputs; low-level `rmsnorm_cuda` requires contiguous inputs | +| `weight` | `[H]` | matches `x` | zero-centred (upstream inits to zeros); CUDA wrapper copies non-contiguous weights; low-level `rmsnorm_cuda` requires contiguous weights | +| `eps` | scalar | float | inside the sqrt; `1e-6` for Qwen3-Next | + +## Dispatch Behavior + +Not registered on this branch: there is no `qwen3_next_rms_norm` gtest spec and no +registry entry yet. Both arrive with the gated-norm PR. Until then, construct the +ops directly as in "Entry Point". The CUDA op validates the compiled symbols in +`__init__`, so on a build without the extension construction raises instead of +handing out an op that fails at call time. + +## Accuracy + +Claim levels: + +- **L1** (prefix-slice and concurrency) for `Qwen3NextRMSNormCudaOp`; padding, + packing and order are not tested for the CUDA op. +- **L1** (slice, concurrency and padding) for `Qwen3NextRMSNormOp`. +- **L0** for `Qwen3NextRMSNormOp` only (CPU, fp32); no test repeats the CUDA op. +- L2 is not claimed. + +The PyTorch reference uses the repo's fixed 32-wide chunked sum +(`shape_invariant_rstd`, introduced in `8ed1693` as a device-agnostic reference). +On B200, over 20 seeds at `H=2048` in bf16, a plain `mean(-1)` broke slice +invariance on 1 of 20 seeds (slices `x[3:5]` and `x[:1]` of 64 rows; which one +failed was not recorded), and the chunked reduction broke on none. That is a recorded +observation, not an assertion. + +The CUDA kernel has its own fixed-order reduction (per-thread strided partial sums, +then `block_reduce_sum`). It is not bitwise equal to the PyTorch reference. + +Neither is bitwise equal to vLLM. In one-off probes (not committed checks), the +reference differed from every vLLM path tried (eager, inductor-compiled, and eager +under `VLLM_BATCH_INVARIANT=1`) on 36–39 of 40 seeds, `max|diff| <= 1.56e-2` +(bf16, `H=2048`, 512 rows). The difference is attributed to reduction order, but that +has not been isolated. Only the no-residual call was compared; vLLM's +`fused_add_rms_norm` path (every decoder norm except layer 0's input norm) has no +reference here. + +The in-kernel offset is exact, not an approximation: `weight_offset=1.0` is bitwise +equal to passing an explicit fp32 `1 + w` weight, and differs from a bf16-folded +`1 + w`, both asserted. + +Accuracy tests resolve their tolerances from `tolerance_contract.json` +(`forward_accuracy`, `reduction` × dtype). The bounds in +`tests/models/qwen3_next/check_qwen3_next_norm_providers.py` are provider-gap bounds, not contract +thresholds. + +## Performance Notes + +The CUDA path reuses the existing `rmsnorm_fwd_kernel` reduction +(`block_reduce_sum` over `choose_threads(H)`), so the offset costs one fp32 add per +element and no extra memory traffic. + +## Evidence + +![zero-centred RMSNorm vs existing implementations on B200](../usage/evidence/qwen3-next-rms-norm-b200/figure.png) + +[`report.json`](../usage/evidence/qwen3-next-rms-norm-b200/report.json) was written by +`tools/validation/models/qwen3_next_norm_evidence.py` from a clean tree at `3d0bae7`, on an otherwise idle +B200 (torch 2.13.0+cu130, transformers 5.17.0, vLLM 0.30.0, FlashInfer 0.6.18). Hidden 2048, +BF16. + +| | rl-kernel CUDA | PyTorch reference | transformers | vLLM `GemmaRMSNorm`* | FlashInfer `gemma_rmsnorm`* | +|---|---|---|---|---|---| +| rows differing alone vs in a batch, forward / `dx` (of 768) | **0 / 0** | 0 / 12 | 2 / 11 | 2 / n/a | 0 / n/a | +| forward, 65536 rows | 425 µs | 2148 µs | 1449 µs | 1449 µs | **90 µs** | +| backward, 65536 rows | 30.9 ms | **3.0 ms** | 3.1 ms | n/a | n/a | + +\* forward only, no backward. + +- **Accuracy is the same for every implementation:** forward max error 1.56e-2 against the + FP64 golden (BF16 output rounding; 99.999% of elements correctly rounded), `dx` and + `dweight` within 2.8e-3 and 2.1e-3 of their maximum. +- **Row invariance:** only this op and FlashInfer give every row the same bits alone and in a + batch; FlashInfer has no backward. transformers and vLLM differ on 2 forward rows in 768. +- **The backward is about 10× slower than transformers at 65536 rows.** `dweight` is folded + over rows in ascending order by `reduce_rows_fp32_left_fold`, one thread per column, so + that the batch `dweight` equals the in-order sum of single-row contributions bitwise (the + gradient-invariance contract's singleton-aggregate check). That serial fold over rows is + the cost; a faster tree fold would not pass that check. + +## Existing implementations: batch invariance of every row, accuracy, gates + +![qwen3_next_rms_norm vs existing implementations](../usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png) + +| Implementation | Batch-invariant | Forward correctly rounded | Worst grad err | Forward / fwd+bwd | C3/C4 gates | +|---|---|---|---|---|---| +| rl-kernel Qwen3NextRMSNormCudaOp | yes | 99.9993% | 2.4e-03 | 426 / 31485 µs | not on this branch | +| rl-kernel PyTorch reference | **no** (49 rows, 94 sub-batches) | 99.9993% | 2.4e-03 | 2146 / 5024 µs | — | +| transformers 5.17.0 Qwen3NextRMSNorm | **no** (90 rows, 142 sub-batches) | 99.9992% | 2.4e-03 | 1449 / 4468 µs | not on this branch | +| Liger 0.8.4 RMSNorm, offset 1, gemma | yes | 99.9993% | 2.4e-03 | 133 / 974 µs | not on this branch | +| FLA 0.5.2 rms_norm, weight passed as 1 + w | yes | 73.0161% | 5.5e-03 | 147 / 1051 µs | not on this branch | +| TE 2.20.2 RMSNorm(zero_centered_gamma) | **no** (175 rows, 245 sub-batches) | 99.9992% | 2.4e-03 | 156 / 779 µs | — | +| FlashInfer 0.6.18.post1 gemma_rmsnorm | yes | 99.9992% | — | 91 / — µs | — | +| vLLM 0.30.0 GemmaRMSNorm.forward_cuda | **no** (16 rows, 26 sub-batches) | 99.9992% | — | 1454 / — µs | — | + +Batch invariance is bitwise and covers three checks: every row computed alone vs inside full +batches of three sizes; the full workload-size batch vs sub-batches that together cover every +row; and a dense batch-size sweep. A "no" counts the rows, sub-batches or sweep cases that +differed. Accuracy is against FP64 at the workload size; latency is the median on an otherwise +idle B200. The C3/C4 column runs this repository's own gate scripts unchanged, with the CUDA +candidate replaced by a subclass of this op whose forward and backward call the other library. +The subclass keeps this op's FP32 `dweight` row contributions, so singleton-aggregate compares +like with like. The gate scripts for this op arrive with #468 (`qwen3_next_norm_manifest.json`); its gate results are in #468's copy of this page. [`qwen3_next_rms_norm.json`](../usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json) +was originally written from a clean tree at `a66493c` by the command below. The +Megatron `52fbcbc` result has been excluded from the report, table and figure: +that revision computes `weight_eff` but uses the original `weight` in its forward +output, so it does not implement the zero-centred operation despite accepting the +flag. The remaining measurements are unchanged; the figure was regenerated from +the corrected report without rerunning benchmarks. The reuse checker now rejects +Megatron implementations that fail a zero-weight probe before benchmarking them. + +Original command: + +```bash +python tools/validation/models/qwen3_next_norm_reuse_check.py --op qwen3_next_rms_norm \ + --out docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json [--megatron-src ] +``` + +## Tests + +```bash +python -m pytest tests/models/qwen3_next/test_qwen3_next_norm.py -v +``` + +## Known Limitations + +- CUDA only; no ROCm, Ascend or Triton backend. +- Not registered as a gtest operator or in the registry on this branch (see + "Dispatch Behavior"). +- Not bitwise against vLLM (see Accuracy). An L2 claim needs a single source of + truth for the forward on both sides, per RFC #428 §0 item 1. +- The gated pair (`Qwen3NextRMSNormGatedOp`, `Qwen3NextRMSNormGatedHFOp`) is + documented with its CUDA kernel in the gated-norm PR, not on this page. +- Measured on sm_100 (B200). Per RFC #428 §2.2 no claim carries across + H100/H200/B100/B200. diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json new file mode 100644 index 000000000..7886c8035 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json @@ -0,0 +1,818 @@ +{ + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": "qwen3_next_rms_norm", + "rl_kernel_commit": "a66493c6d0b49d78aec6df3519b70fb45eace24f", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "cuda": "13.0", + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "liger-kernel": "0.8.4", + "fla-core": "0.5.2", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc" + }, + "quick": false, + "checks": [ + "bi", + "gates", + "perf" + ], + "results": { + "rl_kernel": { + "source": "rl-kernel Qwen3NextRMSNormCudaOp", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999926686286926, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 426.4480024576187, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 31485.487937927246 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "reference": { + "source": "rl-kernel PyTorch reference", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 2 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 7 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 40 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 94 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999927282333374, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 2145.951986312866, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 5024.384021759033 + } + }, + "transformers": { + "source": "transformers 5.17.0 Qwen3NextRMSNorm", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 3 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 10 + }, + "2048": { + "rows": 6144, + "fwd_differ": 16, + "rowgrad_differ": 61 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 26, + "rowgrad_differ": 116 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999924898147583, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 1449.1679668426514, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 4467.727899551392 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "liger": { + "source": "Liger 0.8.4 RMSNorm, offset 1, gemma", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999925494194031, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 133.31200182437897, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 973.6000001430511 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "fla": { + "source": "FLA 0.5.2 rms_norm, weight passed as 1 + w", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.730160653591156, + "fwd_max_abs_err": 0.03517675224226302, + "fwd_us": 146.92799746990204, + "grad_rel_err": { + "x": 0.005450811301370503, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 1050.8000254631042 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "transformer_engine": { + "source": "TE 2.20.2 RMSNorm(zero_centered_gamma)", + "backward": "x-only", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 1, + "rowgrad_differ": 2 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 75, + "rowgrad_differ": 97 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 111, + "rowgrad_differ": 134 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999922513961792, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 156.40000253915787, + "grad_rel_err": { + "x": 0.0023939340856740927 + }, + "fwd_bwd_us": 778.5120010375977 + } + }, + "megatron": { + "unavailable": "Excluded: Megatron 52fbcbc accepts zero_centered_gamma=True but uses the original weight instead of weight_eff in forward." + }, + "flashinfer": { + "source": "FlashInfer 0.6.18.post1 gemma_rmsnorm", + "backward": "none", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999917149543762, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 90.71999788284302 + } + }, + "vllm": { + "source": "vLLM 0.30.0 GemmaRMSNorm.forward_cuda", + "backward": "none", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 16, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 26, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999924898147583, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 1454.3840289115906 + } + } + }, + "evidence_corrections": [ + "Removed the invalid Megatron 52fbcbc zero-centered comparison and regenerated the figure. All other measurements and original run provenance are unchanged; benchmarks were not rerun." + ] +} diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png new file mode 100644 index 000000000..46e954979 Binary files /dev/null and b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png b/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png new file mode 100644 index 000000000..a408aff9c Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json b/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json new file mode 100644 index 000000000..0fc8452ae --- /dev/null +++ b/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json @@ -0,0 +1,225 @@ +{ + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": "3d0bae7174c09a133ea1288d820b63e4a7db8f57", + "git_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer": "0.6.18.post1" + }, + "hidden": 2048, + "eps": 1e-06, + "dtype": "bfloat16", + "ops": { + "zero_centred_rmsnorm": { + "accuracy": { + "257": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999886155128479 + } + }, + "4096": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999991774559021 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999927282333374 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999992847442627 + } + } + }, + "row_invariance": { + "rl-kernel CUDA": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "rl-kernel PyTorch reference": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 12, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "transformers Qwen3NextRMSNorm": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": 11, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "vLLM GemmaRMSNorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + } + }, + "latency": { + "rl-kernel CUDA": { + "1024": { + "forward_us": 24.70399998128414, + "backward_us": 568.7519907951355 + }, + "4096": { + "forward_us": 39.16800022125244, + "backward_us": 1259.3119740486145 + }, + "16384": { + "forward_us": 122.36800044775009, + "backward_us": 7879.024028778076 + }, + "65536": { + "forward_us": 425.80799758434296, + "backward_us": 30937.551498413086 + } + }, + "rl-kernel PyTorch reference": { + "1024": { + "forward_us": 70.54400071501732, + "backward_us": 644.7039842605591 + }, + "4096": { + "forward_us": 155.90400248765945, + "backward_us": 624.2719888687134 + }, + "16384": { + "forward_us": 577.5039792060852, + "backward_us": 986.7520034313202 + }, + "65536": { + "forward_us": 2148.0319499969482, + "backward_us": 2990.224003791809 + } + }, + "transformers Qwen3NextRMSNorm": { + "1024": { + "forward_us": 69.34399902820587, + "backward_us": 608.5599958896637 + }, + "4096": { + "forward_us": 117.3119992017746, + "backward_us": 606.3359975814819 + }, + "16384": { + "forward_us": 407.50400722026825, + "backward_us": 1027.888000011444 + }, + "65536": { + "forward_us": 1449.4240283966064, + "backward_us": 3144.6080207824707 + } + }, + "vLLM GemmaRMSNorm (forward only)": { + "1024": { + "forward_us": 65.61600044369698 + }, + "4096": { + "forward_us": 118.40000003576279 + }, + "16384": { + "forward_us": 405.90400993824005 + }, + "65536": { + "forward_us": 1449.6000409126282 + } + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "1024": { + "forward_us": 14.448000118136406 + }, + "4096": { + "forward_us": 16.24000072479248 + }, + "16384": { + "forward_us": 32.207999378442764 + }, + "65536": { + "forward_us": 90.55999666452408 + } + } + } + } + } +} diff --git a/rl_engine/_C.pyi b/rl_engine/_C.pyi index fb3c4379c..225f997b6 100644 --- a/rl_engine/_C.pyi +++ b/rl_engine/_C.pyi @@ -256,16 +256,21 @@ def swiglu_backward( gate: torch.Tensor, up: torch.Tensor, ) -> list[torch.Tensor]: ... + +rmsnorm_api_version: int + def rmsnorm_forward( x: torch.Tensor, weight: torch.Tensor, eps: float, + weight_offset: float = ..., ) -> list[torch.Tensor]: ... def rmsnorm_backward_dx( dy: torch.Tensor, x: torch.Tensor, weight: torch.Tensor, rstd: torch.Tensor, + weight_offset: float = ..., ) -> torch.Tensor: ... def rmsnorm_backward_dw( dy: torch.Tensor, diff --git a/rl_engine/backends/cuda/norm/rmsnorm.py b/rl_engine/backends/cuda/norm/rmsnorm.py index ec3c65f77..a64b59fde 100644 --- a/rl_engine/backends/cuda/norm/rmsnorm.py +++ b/rl_engine/backends/cuda/norm/rmsnorm.py @@ -4,9 +4,11 @@ from rl_engine.ops.autograd.backward_runtime import record_backward from rl_engine.ops.autograd.vjp_fp32 import reduce_rows_fp32, rmsnorm_dweight_rows_fp32 +_RMSNORM_API_VERSION = 2 -def _require_cuda_symbols(what: str, *names: str) -> None: - """Raise when the compiled kernels backing ``what`` are missing. + +def _require_cuda_rmsnorm() -> None: + """Raise when the compiled RMSNorm bindings are missing or incompatible. The registry treats a backend whose construction raises as unavailable and falls through to the next candidate, so calling this from ``__init__`` is @@ -15,13 +17,21 @@ def _require_cuda_symbols(what: str, *names: str) -> None: activation ops. """ if not _EXT_AVAILABLE or _C is None: - raise RuntimeError(f"{what} requires the compiled rl_engine._C extension.") + raise RuntimeError("CUDA RMSNorm requires the compiled rl_engine._C extension.") + names = ("rmsnorm_forward", "rmsnorm_backward_dx") missing = [name for name in names if not hasattr(_C, name)] if missing: raise RuntimeError( - f"{what} symbols ({', '.join(missing)}) are not compiled into _C. " + f"CUDA RMSNorm symbols ({', '.join(missing)}) are not compiled into _C. " "Rebuild the extension with csrc/cuda/norm/rmsnorm.cu." ) + api_version = getattr(_C, "rmsnorm_api_version", None) + if api_version != _RMSNORM_API_VERSION: + raise RuntimeError( + f"CUDA RMSNorm requires rl_engine._C RMSNorm API version {_RMSNORM_API_VERSION} " + f"(loaded {api_version!r}). Rebuild the extension with csrc/cuda/norm/rmsnorm.cu " + "for weight_offset support." + ) class RMSNormCuda(torch.autograd.Function): @@ -30,20 +40,23 @@ class RMSNormCuda(torch.autograd.Function): """ @staticmethod - def forward(ctx, x, weight, mask=None, eps=1e-6): + def forward(ctx, x, weight, mask=None, eps=1e-6, weight_offset=0.0): """ Forward: - y = x * rsqrt(mean(x^2) + eps) * weight + y = x * rsqrt(mean(x^2) + eps) * (weight_offset + weight) Input: x: [T, H], fp16/bf16/fp32 CUDA tensor weight: [H], fp16/bf16/fp32 CUDA tensor mask: [T], bool CUDA tensor eps: float + weight_offset: float, added to weight in fp32 inside the kernel. + 1.0 selects the zero-centred (1 + w) convention. Output: y: [T, H] """ + _require_cuda_rmsnorm() assert x.is_cuda, "x must be CUDA tensor" assert weight.is_cuda, "weight must be CUDA tensor" assert x.is_contiguous(), "x must be contiguous" @@ -51,10 +64,6 @@ def forward(ctx, x, weight, mask=None, eps=1e-6): assert x.dim() == 2, "x must be [T, H]" assert weight.dim() == 1, "weight must be [H]" assert x.shape[1] == weight.shape[0], "hidden size mismatch" - assert _EXT_AVAILABLE and hasattr( - _C, "rmsnorm_forward" - ), "RMSNorm CUDA extension is unavailable. Please rebuild with rmsnorm.cu." - if mask is None: mask = torch.ones((x.shape[0],), device=x.device, dtype=torch.bool) else: @@ -64,10 +73,11 @@ def forward(ctx, x, weight, mask=None, eps=1e-6): assert mask.dim() == 1, "mask must be [T]" assert mask.shape[0] == x.shape[0], "mask length mismatch" - y, rstd = _C.rmsnorm_forward(x, weight, float(eps)) + y, rstd = _C.rmsnorm_forward(x, weight, float(eps), float(weight_offset)) ctx.save_for_backward(x, weight, rstd, mask) ctx.eps = eps + ctx.weight_offset = float(weight_offset) return y @@ -81,7 +91,7 @@ def backward(ctx, grad_out): x, weight, rstd, mask = ctx.saved_tensors dy = grad_out.contiguous() - dx = _C.rmsnorm_backward_dx(dy, x, weight, rstd) + dx = _C.rmsnorm_backward_dx(dy, x, weight, rstd, ctx.weight_offset) # The shape-independent FP32 left fold preserves the C2 Batch/Chunk # reduction order while the CUDA reducer executes it in one launch. @@ -101,16 +111,18 @@ def backward(ctx, grad_out): family="cuda", ) - return dx, dw, None, None + # dw is unchanged by the offset: d/dw (offset + w) == d/dw w. + return dx, dw, None, None, None -def rmsnorm_cuda(x, weight, eps=1e-6, mask=None): +def rmsnorm_cuda(x, weight, eps=1e-6, mask=None, weight_offset=0.0): """ use: y = rmsnorm_cuda(x, weight) y = rmsnorm_cuda(x, weight, mask=mask) + y = rmsnorm_cuda(x, weight, weight_offset=1.0) # zero-centred weight """ - return RMSNormCuda.apply(x, weight, mask, eps) + return RMSNormCuda.apply(x, weight, mask, eps, weight_offset) class RMSNormCudaOp: @@ -119,11 +131,11 @@ class RMSNormCudaOp: backward_impl = "cuda_rmsnorm_dx_declared_fp32_rowfold_dw" def __init__(self) -> None: - _require_cuda_symbols( - "CUDA RMSNorm", - "rmsnorm_forward", - "rmsnorm_backward_dx", - ) + _require_cuda_rmsnorm() + + #: Added to the weight in fp32 inside the kernel. Subclasses override it; + #: 0.0 is the plain convention. + weight_offset = 0.0 def __call__(self, x, weight, *, eps=1e-6): return self.forward(x, weight, eps=eps) @@ -131,12 +143,29 @@ def __call__(self, x, weight, *, eps=1e-6): def forward(self, x, weight, *, eps=1e-6): hidden = x.shape[-1] x_2d = x.contiguous().view(-1, hidden) - y_2d = rmsnorm_cuda(x_2d, weight.contiguous(), eps=eps) + y_2d = rmsnorm_cuda(x_2d, weight.contiguous(), eps=eps, weight_offset=self.weight_offset) return y_2d.view_as(x) def parameter_vjp_contributions_fp32(self, *, x, weight, grad_output, eps=1e-6): - del weight - x32 = x.float() - rstd = torch.rsqrt(x32.square().mean(dim=-1) + float(eps)) - rows = grad_output.float() * x32 * rstd.unsqueeze(-1) + hidden = x.shape[-1] + # Only `rstd` is used, and it does not depend on the offset; the offset is + # passed so this is the same call the forward makes, not because it matters. + _, rstd = _C.rmsnorm_forward( + x.contiguous().reshape(-1, hidden), + weight.contiguous(), + float(eps), + float(self.weight_offset), + ) + rows = rmsnorm_dweight_rows_fp32(x, grad_output, rstd=rstd.reshape(x.shape[:-1])) return {"weight": rows} + + +class Qwen3NextRMSNormCudaOp(RMSNormCudaOp): + """Zero-centred CUDA RMSNorm: ``y = x * rstd * (1 + weight)``. + + The decoder and final norms of Qwen3-Next (and Gemma) store a zero-centred + weight. The ``+1`` is applied inside the kernel after the fp32 upcast, so it + is never rounded through the low-precision weight dtype. + """ + + weight_offset = 1.0 diff --git a/rl_engine/reference/norm/qwen3_next_rms_norm.py b/rl_engine/reference/norm/qwen3_next_rms_norm.py new file mode 100644 index 000000000..adb024bcc --- /dev/null +++ b/rl_engine/reference/norm/qwen3_next_rms_norm.py @@ -0,0 +1,158 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Qwen3-Next RMSNorm references (WS1 ground truth for RFC #428 C1). + +Qwen3-Next ships two RMSNorm conventions. They differ in how the weight is applied +and in where the dtype casts sit, so they are separate operators rather than a flag: + +``Qwen3NextRMSNorm`` (decoder / final norm) + Normalize in fp32, scale by ``(1 + weight)`` in fp32, cast once at the end. The + stored weight is zero-centred, so the ``1 +`` must be applied after the fp32 + upcast; folding it into a bf16 weight first rounds the offset. + +``Qwen3NextRMSNormGated`` (inside the Gated DeltaNet block) + Normalize in fp32, scale by a plain weight, then gate by ``silu(gate)``. + vLLM's ``RMSNormGated`` multiplies the weight in fp32; transformers casts the + normalized value back to the input dtype first. ``Qwen3NextRMSNormGatedOp`` + follows vLLM; ``Qwen3NextRMSNormGatedHFOp`` keeps the transformers convention + as a witness, and ``test_gated_conventions_diverge_in_low_precision`` pins that + the two differ in bf16 and agree bitwise in fp32. + +All three reuse :func:`shape_invariant_rstd`, a fixed-order reduction. They +reproduce the weight convention and cast order, not vLLM's reduction tree, and are +not bitwise equal to any vLLM path probed so far. Claim levels, measurements and +limitations are in ``docs/operators/qwen3-next-rms-norm.md``. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.reference.norm.rms_norm import ( + NativeRMSNormOp, + check_norm_weight, + shape_invariant_rstd, +) + +__all__ = [ + "Qwen3NextRMSNormOp", + "Qwen3NextRMSNormGatedOp", + "Qwen3NextRMSNormGatedHFOp", +] + + +class Qwen3NextRMSNormOp(NativeRMSNormOp): + """Zero-centred RMSNorm: ``out = x * rstd * (1 + weight)``. + + Only the weight convention differs from :class:`NativeRMSNormOp`, so that is + all this overrides. The base applies the offset in fp32, after the upcast, + which is what ``transformers`` and vLLM both do -- folding ``1 +`` into a + bf16 weight beforehand would round the offset away. + """ + + weight_offset = 1.0 + + +class Qwen3NextRMSNormGatedOp: + """Gated RMSNorm used by the Gated DeltaNet block (vLLM/strict convention). + + ``out = (x * rstd * weight) * silu(gate)``, with every multiply in fp32 and + a single cast on the way out. This is the convention vLLM's ``RMSNormGated`` + uses with ``norm_before_gate=True``. Which gated convention is the strict + default is still open; see the operator page. + + Not a subclass of the plain op: it takes an extra tensor and its epilogue + differs, so it is not a drop-in substitute for one. + + The weight is plain, NOT zero-centred -- upstream initializes it to ones. + """ + + def __call__( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + return self.forward(x, weight, gate, eps=eps) + + def forward( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + return self._rms_norm_gated(x, weight, gate, eps=eps, output_dtype=x.dtype) + + def forward_fp32( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + """Ground-truth: fp32 in, fp32 out.""" + return self._rms_norm_gated(x, weight, gate, eps=eps, output_dtype=torch.float32) + + # ------------------------------------------------------------------ # + # Shared by both conventions; only `_scale_by_weight` differs. + # ------------------------------------------------------------------ # + @staticmethod + def _normalized( + x: torch.Tensor, weight: torch.Tensor, gate: torch.Tensor, *, eps: float + ) -> torch.Tensor: + check_norm_weight(x, weight) + if gate.shape != x.shape: + raise ValueError( + f"gate must match x, got tuple(gate.shape)={tuple(gate.shape)} " + f"vs tuple(x.shape)={tuple(x.shape)}" + ) + x_f = x.float() + rstd = shape_invariant_rstd(x_f, float(eps)).unsqueeze(-1) + return x_f * rstd + + @staticmethod + def _scale_by_weight( + normed: torch.Tensor, weight: torch.Tensor, input_dtype: torch.dtype + ) -> torch.Tensor: + """vLLM: the weight multiply stays in fp32.""" + del input_dtype + return normed * weight.float() + + @classmethod + def _rms_norm_gated( + cls, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float, + output_dtype: torch.dtype, + ) -> torch.Tensor: + normed = cls._normalized(x, weight, gate, eps=eps) + scaled = cls._scale_by_weight(normed, weight, x.dtype) + # The gate activation is evaluated in fp32 and promotes the product. + gated = scaled * F.silu(gate.float()) + return gated.to(output_dtype) + + +class Qwen3NextRMSNormGatedHFOp(Qwen3NextRMSNormGatedOp): + """``transformers`` gated convention: weight multiply in the input dtype. + + Kept so the HF-vs-vLLM divergence documented in the module docstring stays + covered by a test rather than discovered in a drift report. Do NOT use this + for an L2 claim against vLLM rollout. + """ + + @staticmethod + def _scale_by_weight( + normed: torch.Tensor, weight: torch.Tensor, input_dtype: torch.dtype + ) -> torch.Tensor: + """transformers: round-trip through the input dtype first.""" + return weight * normed.to(input_dtype) diff --git a/rl_engine/reference/norm/rms_norm.py b/rl_engine/reference/norm/rms_norm.py index 76dc13b49..e08738a2e 100644 --- a/rl_engine/reference/norm/rms_norm.py +++ b/rl_engine/reference/norm/rms_norm.py @@ -86,6 +86,15 @@ def strict_add_rms_norm( return _strict_add_rms_norm(x, residual, weight, eps) +def check_norm_weight(x: torch.Tensor, weight: torch.Tensor) -> None: + """Shared shape guard for the RMSNorm family (plain, zero-centred, gated).""" + if weight.dim() != 1 or weight.shape[0] != x.shape[-1]: + raise ValueError( + f"weight must be 1-D of size x.shape[-1]={x.shape[-1]}, " + f"got tuple(weight.shape)={tuple(weight.shape)}" + ) + + def shape_invariant_rstd(x_f: torch.Tensor, eps: float) -> torch.Tensor: """Shape-invariant per-row rstd (the shared RMSNorm statistic). @@ -113,6 +122,10 @@ class NativeRMSNormOp: out = x * rsqrt(mean(x^2, dim=-1) + eps) * weight """ + #: Added to the weight in fp32 before it scales the normalized value. + #: Subclasses set 1.0 for the zero-centred ``(1 + w)`` convention. + weight_offset = 0.0 + def __init__(self) -> None: pass @@ -151,21 +164,24 @@ def forward_fp32( # ------------------------------------------------------------------ # # Helpers # ------------------------------------------------------------------ # - @staticmethod + @classmethod def _rms_norm( + cls, x: torch.Tensor, weight: torch.Tensor, *, eps: float, output_dtype: torch.dtype, ) -> torch.Tensor: - if weight.dim() != 1 or weight.shape[0] != x.shape[-1]: - raise ValueError( - f"weight must be 1-D of size x.shape[-1]={x.shape[-1]}, " - f"got tuple(weight.shape)={tuple(weight.shape)}" - ) + check_norm_weight(x, weight) x_f = x.float() rstd = shape_invariant_rstd(x_f, float(eps)).unsqueeze(-1) normed = x_f * rstd - out = normed * weight.float() + scale = weight.float() + # Guarded rather than unconditional: `0.0 + w` rewrites -0.0 to +0.0, + # which torch.equal does not notice but a bitwise comparison does. The + # plain path must stay bit-for-bit what it was. + if cls.weight_offset: + scale = cls.weight_offset + scale + out = normed * scale return out.to(output_dtype) diff --git a/tests/models/qwen3_next/__init__.py b/tests/models/qwen3_next/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/tests/models/qwen3_next/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/tests/models/qwen3_next/check_qwen3_next_norm_providers.py b/tests/models/qwen3_next/check_qwen3_next_norm_providers.py new file mode 100644 index 000000000..f4d25d2d8 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_norm_providers.py @@ -0,0 +1,239 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""How far apart vLLM's eager gated-RMSNorm paths are, and ours from each. + +Named ``check_`` rather than ``test_`` on purpose, following +``tests/distributed/check_*.py``: this module imports real vLLM, and +The integration isolation test asserts that ``vllm`` is absent +from ``sys.modules`` -- an invariant any collected test importing vLLM would break +for the whole session. Run it explicitly: + + pytest tests/models/qwen3_next/check_qwen3_next_norm_providers.py -v + +RFC #428 section 2.1 defines L2 as "bitwise identical to vLLM rollout". For +Qwen3-Next's GDN gated norm that phrase is not well defined until a single +provider is named, because vLLM ships several and they do not agree bitwise with +each other. + +This module records two things so a vLLM upgrade cannot move them silently: + +1. **Provider facts** -- the GDN decode env defaults, and that ``RMSNormGated`` + has distinct ``forward_native`` and ``forward_cuda`` methods. It does NOT + determine which path vLLM dispatches at runtime: the checks below call each + method directly (eager), and vLLM's default compiled mode traces + ``forward_native`` into an inductor graph instead. +2. **The size of the gap** -- a seed sweep that asserts an upper bound on the + disagreement and on the mismatch rate between the eager paths. The bounds are + provider-gap bounds, not ``tolerance_contract.json`` thresholds, and they + deliberately do NOT assert equality. + +Measured on 2x B200 (sm_100, torch 2.13.0+cu130, vllm 0.30.0), bf16, +``head_v_dim=128``, 512 rows, 40 seeds: + +=============================== ================== ================ +comparison seeds not bitwise worst max|diff| +=============================== ================== ================ +ours vs ``forward_native`` 6 / 40 1.56e-2 +ours vs ``forward_cuda`` 18 / 40 3.91e-3 +``forward_native`` vs ``cuda`` 21 / 40 1.56e-2 +=============================== ================== ================ +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormGatedOp + +# vLLM is imported inside the checks, never at module scope, so that merely +# collecting this file does not pull it into sys.modules. + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") + +# Qwen3-Next-80B-A3B-Instruct: linear_value_head_dim / rms_norm_eps. +_HEAD_V_DIM = 128 +_EPS = 1e-6 +_ROWS = 512 +_SEEDS = 40 + +# The gated norm is constructed by vLLM's GDN block with exactly these settings; +# see vllm/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py. +_NORM_BEFORE_GATE = True +_GROUP_SIZE = None +_ACTIVATION = "silu" + + +@pytest.fixture(scope="module") +def vllm_config_ctx(): + """`RMSNormGated` is a CustomOp and refuses to build outside a config context.""" + pytest.importorskip("vllm", reason="vLLM is required to identify the provider") + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + yield + + +def _make_norm(): + from vllm.model_executor.layers.layernorm import RMSNormGated + + return RMSNormGated( + _HEAD_V_DIM, + eps=_EPS, + group_size=_GROUP_SIZE, + norm_before_gate=_NORM_BEFORE_GATE, + activation=_ACTIVATION, + ) + + +def _inputs(seed: int, dtype: torch.dtype, rows: int = _ROWS): + g = torch.Generator(device="cuda").manual_seed(seed) + x = torch.randn(rows, _HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + gate = torch.randn(rows, _HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + weight = torch.randn(_HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + return x, gate, weight + + +def _disagreement(a: torch.Tensor, b: torch.Tensor) -> tuple[float, float]: + """(worst absolute difference, fraction of elements that differ bitwise).""" + bits_a = a.float().view(torch.int32) + bits_b = b.float().view(torch.int32) + mismatch = int((bits_a != bits_b).sum()) + worst = (a.float() - b.float()).abs().max().item() + return worst, mismatch / a.numel() + + +# --------------------------------------------------------------------------- # +# 1. Provider facts -- recorded, not a runtime dispatch check +# --------------------------------------------------------------------------- # +def test_custom_op_has_distinct_native_and_cuda_paths(vllm_config_ctx): + """The two methods are distinct. Which one runs depends on the vLLM mode: eager + (``custom_ops="all"``) dispatches ``forward_cuda``; the default compiled mode + traces ``forward_native``. This test does not check that choice.""" + norm = _make_norm() + assert type(norm).forward_cuda is not type(norm).forward_native + + +def test_gdn_decode_provider_env_defaults_are_recorded(): + """Record the GDN decode env defaults. + + This records the defaults only; it does not assert which kernel runs. For + Qwen3-Next the ``VLLM_GDN_DECODE_KERNEL="cuda"`` default does not take effect: + vLLM builds its GDN layers with ``gqa_interleaved_layout=True`` and falls back + to the Triton decode kernel. A change of either default still fails here, as a + prompt to re-derive the decode path. + """ + pytest.importorskip("vllm", reason="vLLM is required to identify the provider") + import vllm.envs as envs + + observed = { + "VLLM_GDN_DECODE_KERNEL": envs.VLLM_GDN_DECODE_KERNEL.strip().lower(), + "VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE": bool( + envs.VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE + ), + } + assert observed == { + "VLLM_GDN_DECODE_KERNEL": "cuda", + "VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE": True, + }, ( + f"GDN decode provider defaults changed: {observed}. Re-derive which kernel " + "the rollout decode path takes before relying on any exactness claim." + ) + + +# --------------------------------------------------------------------------- # +# 2. The gap, bounded -- never asserted to be zero +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize( + "dtype, max_abs, max_mismatch_rate", + [ + # Provider-gap bounds, not contract thresholds. bf16 rounding absorbs most + # of the reduction-tree difference, so few elements move. + (torch.bfloat16, 2e-2, 0.05), + # In fp32 about 36% of elements differ, so the rate is left unbounded and + # only the magnitude is bounded (no per-element ULP bound is asserted). + (torch.float32, 1e-5, 1.0), + ], +) +def test_vllm_paths_disagree_only_within_bounds(vllm_config_ctx, dtype, max_abs, max_mismatch_rate): + """vLLM's own two paths differ; bound how much. + + A growing magnitude means the reduction trees have diverged semantically, + which would invalidate treating either as the reference. A high mismatch + *rate* at ULP magnitude does not -- it just means the tree shapes differ. + """ + norm = _make_norm().to("cuda", dtype) + worst_abs, worst_rate, differing = 0.0, 0.0, 0 + for seed in range(_SEEDS): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + native = norm.forward_native(x, gate) + cuda = norm.forward_cuda(x, gate) + abs_d, rate = _disagreement(native, cuda) + worst_abs, worst_rate = max(worst_abs, abs_d), max(worst_rate, rate) + differing += int(rate > 0.0) + + assert worst_abs <= max_abs, ( + f"vLLM forward_native vs forward_cuda worst |diff| {worst_abs:.3e} exceeds " + f"{max_abs:.3e} over {_SEEDS} seeds ({differing} seeds differ)" + ) + assert worst_rate <= max_mismatch_rate + + +@pytest.mark.parametrize("path", ["forward_native", "forward_cuda"]) +def test_ours_tracks_each_vllm_path_within_bounds(vllm_config_ctx, path): + """Our strict op follows vLLM's convention; bound the residual. + + Not an equality assertion. We reproduce the fp32 weight multiply and the + single trailing cast, but our reduction is the repo's fixed 32-wide chunked + sum rather than whatever tree the provider uses, so a few elements straddle + a rounding boundary. + """ + dtype = torch.bfloat16 + ours = Qwen3NextRMSNormGatedOp() + norm = _make_norm().to("cuda", dtype) + + worst_abs, worst_rate = 0.0, 0.0 + for seed in range(_SEEDS): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + reference = getattr(norm, path)(x, gate) + abs_d, rate = _disagreement(ours.forward(x, weight, gate), reference) + worst_abs, worst_rate = max(worst_abs, abs_d), max(worst_rate, rate) + + # Provider-gap bounds, not contract thresholds. + assert worst_abs <= 2e-2, f"worst |diff| vs {path} was {worst_abs:.3e}" + assert worst_rate <= 0.05, f"mismatch rate vs {path} was {worst_rate:.3%}" + + +# --------------------------------------------------------------------------- # +# 3. Batch invariance -- the property we DO claim +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("path", ["forward_native", "forward_cuda"]) +def test_vllm_gated_norm_is_batch_invariant(vllm_config_ctx, path): + """If this ever fails, vLLM stops being a coherent L2 target. + + Kept as a guard rather than a claim about our code: a provider whose output + depends on unrelated rows cannot anchor a bitwise contract. + """ + dtype = torch.bfloat16 + norm = _make_norm().to("cuda", dtype) + for seed in range(8): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + full = getattr(norm, path)(x, gate) + for n in (1, 2, 8, 16, 32, 48, 64, 256): + sliced = getattr(norm, path)(x[:n], gate[:n]) + assert torch.equal(sliced, full[:n]), f"{path} seed={seed} n={n}" + + +def test_our_gated_op_is_batch_invariant(): + """The L1 claim for this operator, bitwise.""" + dtype = torch.bfloat16 + ours = Qwen3NextRMSNormGatedOp() + for seed in range(8): + x, gate, weight = _inputs(seed, dtype) + full = ours.forward(x, weight, gate) + for n in (1, 2, 8, 16, 32, 48, 64, 256): + assert torch.equal(ours.forward(x[:n], weight, gate[:n]), full[:n]) diff --git a/tests/models/qwen3_next/test_qwen3_next_norm.py b/tests/models/qwen3_next/test_qwen3_next_norm.py new file mode 100644 index 000000000..5c9d5feff --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_norm.py @@ -0,0 +1,555 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""WS1 C1 (RFC #428): Qwen3-Next RMSNorm conventions and batch invariance. + +Claim levels exercised here (RFC #428 section 2.1): + * L0 repeatable -- identical inputs reproduce bitwise identical outputs. + * L1 batch-invariant -- a row is unaffected by slicing, concurrency and (for the + PyTorch reference) padding; packing and order are not exercised. + +L2 (train-rollout exact against vLLM) is NOT claimed by this file; it needs the +rollout engine on the other side. + +The reduction order here is the repo's fixed 32-wide chunked reduction, which +differs from upstream's ``mean(-1)``. What is reproduced exactly is the weight +convention and the cast order, so comparisons against the upstream formula are +tolerance based while the invariance assertions are bitwise. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.contracts.numerical import load_contract, resolve_tolerance +from rl_engine.reference.norm.qwen3_next_rms_norm import ( + Qwen3NextRMSNormGatedHFOp, + Qwen3NextRMSNormGatedOp, + Qwen3NextRMSNormOp, +) + +# Qwen3-Next-80B-A3B-Instruct config.json +_HIDDEN = 2048 # hidden_size +_HEAD_V_DIM = 128 # linear_value_head_dim -- the gated norm width +_EPS = 1e-6 # rms_norm_eps + +_CONTRACT = load_contract() + + +def _forward_tol(dtype: torch.dtype) -> dict[str, float]: + """C1 forward_accuracy row for the ``reduction`` op class -- no private thresholds.""" + spec = resolve_tolerance( + _CONTRACT, judgment="forward_accuracy", op_class="reduction", dtype=dtype + ) + return {"atol": spec.atol, "rtol": spec.rtol} + + +def _rand(shape, seed): + g = torch.Generator().manual_seed(seed) + return torch.randn(*shape, generator=g, dtype=torch.float32) + + +# --------------------------------------------------------------------------- # +# Upstream formulas, transcribed from transformers/models/qwen3_next. +# These pin the exact semantics; they are deliberately verbatim. +# --------------------------------------------------------------------------- # +def _hf_rms_norm(x, weight, eps=_EPS): + output = x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + eps) + output = output * (1.0 + weight.float()) + return output.type_as(x) + + +def _vllm_rms_norm_gated(x, weight, gate, eps=_EPS): + """vLLM ``RMSNormGated.forward_static`` with ``norm_before_gate=True``.""" + orig_dtype = x.dtype + x, weight, z = x.float(), weight.float(), gate.float() + variance = x.pow(2).mean(dim=-1, keepdim=True) + out = (x * torch.rsqrt(variance + eps)) * weight + out = out * F.silu(z) + return out.to(orig_dtype) + + +def _hf_rms_norm_gated(x, weight, gate, eps=_EPS): + input_dtype = x.dtype + h = x.to(torch.float32) + variance = h.pow(2).mean(-1, keepdim=True) + h = h * torch.rsqrt(variance + eps) + h = weight * h.to(input_dtype) + h = h * F.silu(gate.to(torch.float32)) + return h.to(input_dtype) + + +# --------------------------------------------------------------------------- # +# 1. The zero-centred (1 + w) convention +# --------------------------------------------------------------------------- # +def test_zero_weight_is_identity_scaling(): + """weight == 0 must leave the normalized value untouched, bitwise. + + This is what separates the zero-centred convention from the plain one: a + plain RMSNorm returns all zeros for a zero weight, this one returns the + bare normalized value. + """ + from rl_engine.reference.norm.rms_norm import NativeRMSNormOp, shape_invariant_rstd + + x = _rand((4, _HIDDEN), seed=0) + zeros = torch.zeros(_HIDDEN) + + x_f = x.float() + bare = x_f * shape_invariant_rstd(x_f, _EPS).unsqueeze(-1) + assert torch.equal(Qwen3NextRMSNormOp().forward_fp32(x, zeros, eps=_EPS), bare) + + # ... and the plain convention really does differ here. + assert torch.equal(NativeRMSNormOp().forward_fp32(x, zeros, eps=_EPS), torch.zeros_like(bare)) + + +def test_weight_offset_is_one(): + """w and w+1 scaling relationship: out(w) == norm * (1 + w).""" + op = Qwen3NextRMSNormOp() + x = _rand((2, _HEAD_V_DIM), seed=1) + unit = op.forward_fp32(x, torch.zeros(_HEAD_V_DIM)) # scale = 1 + doubled = op.forward_fp32(x, torch.ones(_HEAD_V_DIM)) # scale = 2 + torch.testing.assert_close(doubled, 2.0 * unit, atol=1e-6, rtol=1e-6) + + +def test_offset_applied_in_fp32_not_folded_into_bf16(): + """The 1 + w offset must not be pre-rounded through bf16. + + A weight one bf16 ULP below zero stays distinguishable from exactly zero + once the offset is added in fp32; folding (1 + w) into bf16 first would + collapse both to 1.0 and lose the difference. + """ + op = Qwen3NextRMSNormOp() + x = _rand((2, _HEAD_V_DIM), seed=2) + tiny = torch.full((_HEAD_V_DIM,), -(2**-9), dtype=torch.bfloat16) + folded = (1.0 + tiny.float()).bfloat16() # the wrong way + assert not torch.equal( + op.forward_fp32(x, tiny), + op.forward_fp32(x, folded - 1.0), + ) + + +# --------------------------------------------------------------------------- # +# 2. L1 -- batch invariance, bitwise +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("hidden", [_HIDDEN, _HEAD_V_DIM]) +def test_batch_invariance_slice(hidden): + op = Qwen3NextRMSNormOp() + w, x = _rand((hidden,), seed=3), _rand((8, 32, hidden), seed=4) + full = op.forward_fp32(x, w) + assert torch.equal(op.forward_fp32(x[:1], w), full[:1]) + assert torch.equal(op.forward_fp32(x[3:5], w), full[3:5]) + + +@pytest.mark.parametrize("batch", [1, 2, 8, 16, 32, 48, 64]) +def test_batch_invariance_across_concurrency(batch): + """RFC #428 section 10: batch size / concurrency axis 1..64.""" + op = Qwen3NextRMSNormOp() + w = _rand((_HIDDEN,), seed=5) + target = _rand((1, _HIDDEN), seed=6) + others = _rand((batch - 1, _HIDDEN), seed=7) if batch > 1 else None + alone = op.forward_fp32(target, w) + for position in ("first", "last"): + if others is None: + batched, index = target, 0 + elif position == "first": + batched, index = torch.cat([target, others]), 0 + else: + batched, index = torch.cat([others, target]), batch - 1 + assert torch.equal(op.forward_fp32(batched, w)[index : index + 1], alone) + + +def test_batch_invariance_with_padding(): + op = Qwen3NextRMSNormOp() + w = _rand((_HIDDEN,), seed=8) + x = _rand((4, _HIDDEN), seed=9) + padded = torch.cat([x, _rand((6, _HIDDEN), seed=10)], dim=0) + assert torch.equal(op.forward_fp32(padded, w)[:4], op.forward_fp32(x, w)) + + +def test_gated_batch_invariance_slice(): + op = Qwen3NextRMSNormGatedOp() + w = _rand((_HEAD_V_DIM,), seed=11) + x = _rand((8, 32, _HEAD_V_DIM), seed=12) + gate = _rand((8, 32, _HEAD_V_DIM), seed=13) + full = op.forward_fp32(x, w, gate) + assert torch.equal(op.forward_fp32(x[:1], w, gate[:1]), full[:1]) + assert torch.equal(op.forward_fp32(x[3:5], w, gate[3:5]), full[3:5]) + + +# --------------------------------------------------------------------------- # +# 2b. The fixed-order reduction is what buys invariance -- on CUDA too +# --------------------------------------------------------------------------- # +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("hidden", [_HIDDEN, _HEAD_V_DIM]) +def test_chunked_reduction_is_slice_invariant_on_device(hidden): + """The rstd statistic must be slice-invariant on the accelerator, every seed. + + Measured counterpoint on B200/bf16/H=2048: a plain ``mean(-1)`` broke this on + 1 of 20 seeds. We assert only the property we rely on -- asserting that torch's + ``mean`` is broken would be a brittle test of someone else's kernel. + """ + from rl_engine.reference.norm.rms_norm import shape_invariant_rstd + + for seed in range(20): + g = torch.Generator(device="cuda").manual_seed(seed) + x = torch.randn(64, hidden, device="cuda", dtype=torch.bfloat16, generator=g) + full = shape_invariant_rstd(x.float(), _EPS) + assert torch.equal(shape_invariant_rstd(x[3:5].float(), _EPS), full[3:5]), seed + assert torch.equal(shape_invariant_rstd(x[:1].float(), _EPS), full[:1]), seed + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_op_batch_invariance_on_device(): + """L1 for the op itself, on the accelerator, in the model's dtype.""" + op = Qwen3NextRMSNormOp() + g = torch.Generator(device="cuda").manual_seed(0) + x = torch.randn(64, _HIDDEN, device="cuda", dtype=torch.bfloat16, generator=g) + w = torch.randn(_HIDDEN, device="cuda", dtype=torch.bfloat16, generator=g) + full = op.forward(x, w) + for n in (1, 2, 8, 16, 32, 48, 64): + assert torch.equal(op.forward(x[:n], w), full[:n]), n + + +# --------------------------------------------------------------------------- # +# 3. L0 -- repeatable +# --------------------------------------------------------------------------- # +def test_deterministic_repeat(): + op = Qwen3NextRMSNormOp() + x, w = _rand((64, _HIDDEN), seed=14), _rand((_HIDDEN,), seed=15) + first = op.forward_fp32(x, w) + for _ in range(10): + assert torch.equal(op.forward_fp32(x, w), first) + + +# --------------------------------------------------------------------------- # +# 4. Agreement with the upstream formula (tolerance, not bitwise: the reduction +# order differs by design) +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("hidden", [_HIDDEN, _HEAD_V_DIM]) +def test_matches_upstream_formula_fp32(hidden): + op = Qwen3NextRMSNormOp() + x, w = _rand((4, 16, hidden), seed=16), _rand((hidden,), seed=17) + torch.testing.assert_close(op.forward_fp32(x, w), _hf_rms_norm(x, w), atol=1e-6, rtol=1e-6) + + +def test_gated_strict_matches_vllm_formula_fp32(): + """The strict gated op follows vLLM, which is what the L2 claim targets.""" + op = Qwen3NextRMSNormGatedOp() + x = _rand((4, 16, _HEAD_V_DIM), seed=18) + w = _rand((_HEAD_V_DIM,), seed=19) + gate = _rand((4, 16, _HEAD_V_DIM), seed=20) + torch.testing.assert_close( + op.forward_fp32(x, w, gate), + _vllm_rms_norm_gated(x, w, gate), + atol=1e-6, + rtol=1e-6, + ) + + +def test_gated_hf_witness_matches_hf_formula_fp32(): + op = Qwen3NextRMSNormGatedHFOp() + x = _rand((4, 16, _HEAD_V_DIM), seed=18) + w = _rand((_HEAD_V_DIM,), seed=19) + gate = _rand((4, 16, _HEAD_V_DIM), seed=20) + torch.testing.assert_close( + op.forward_fp32(x, w, gate), + _hf_rms_norm_gated(x, w, gate), + atol=1e-6, + rtol=1e-6, + ) + + +def test_gated_conventions_diverge_in_low_precision(): + """Pin the HF-vs-vLLM gated divergence (RFC #428 first-divergence boundary). + + With fp32 inputs the two conventions coincide, because the HF round-trip is + an identity. In bf16 they do not, and the gap is far larger than a ULP: this + is why the convention is part of the operator identity and not a detail. + """ + strict, witness = Qwen3NextRMSNormGatedOp(), Qwen3NextRMSNormGatedHFOp() + x = _rand((64, _HEAD_V_DIM), seed=33).bfloat16() + w = _rand((_HEAD_V_DIM,), seed=34).bfloat16() + gate = _rand((64, _HEAD_V_DIM), seed=35).bfloat16() + + a, b = strict.forward(x, w, gate), witness.forward(x, w, gate) + assert not torch.equal(a, b), "the two gated conventions must be distinguishable" + assert (a.float() - b.float()).abs().max() > 1e-3 + + # ... and in fp32 the round-trip vanishes, so they agree bitwise. + xf, wf, gf = x.float(), w.float(), gate.float() + assert torch.equal(strict.forward(xf, wf, gf), witness.forward(xf, wf, gf)) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_matches_upstream_formula_low_precision(dtype): + op = Qwen3NextRMSNormOp() + x = _rand((4, 16, _HIDDEN), seed=21).to(dtype) + w = _rand((_HIDDEN,), seed=22).to(dtype) + got, ref = op.forward(x, w), _hf_rms_norm(x, w) + assert got.dtype == ref.dtype == dtype + torch.testing.assert_close(got.float(), ref.float(), **_forward_tol(dtype)) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_gated_hf_witness_low_precision(dtype): + """Covers the mid-computation cast back to the input dtype.""" + op = Qwen3NextRMSNormGatedHFOp() + x = _rand((4, 16, _HEAD_V_DIM), seed=23).to(dtype) + w = _rand((_HEAD_V_DIM,), seed=24).to(dtype) + gate = _rand((4, 16, _HEAD_V_DIM), seed=25).to(dtype) + got, ref = op.forward(x, w, gate), _hf_rms_norm_gated(x, w, gate) + assert got.dtype == ref.dtype == dtype + torch.testing.assert_close(got.float(), ref.float(), **_forward_tol(dtype)) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_gated_strict_low_precision(dtype): + op = Qwen3NextRMSNormGatedOp() + x = _rand((4, 16, _HEAD_V_DIM), seed=23).to(dtype) + w = _rand((_HEAD_V_DIM,), seed=24).to(dtype) + gate = _rand((4, 16, _HEAD_V_DIM), seed=25).to(dtype) + got, ref = op.forward(x, w, gate), _vllm_rms_norm_gated(x, w, gate) + assert got.dtype == ref.dtype == dtype + torch.testing.assert_close(got.float(), ref.float(), **_forward_tol(dtype)) + + +# --------------------------------------------------------------------------- # +# 5. dtype paths and guards +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +def test_dtype_paths(dtype): + op = Qwen3NextRMSNormOp() + x = _rand((2, 16, _HIDDEN), seed=26).to(dtype) + w = _rand((_HIDDEN,), seed=27).to(dtype) + assert op.forward(x, w).dtype == dtype + assert op.forward_fp32(x, w).dtype == torch.float32 + + +def test_eps_inside_sqrt(): + op = Qwen3NextRMSNormOp() + out = op.forward_fp32(torch.zeros(1, _HIDDEN), torch.zeros(_HIDDEN)) + assert torch.isfinite(out).all() + assert torch.equal(out, torch.zeros(1, _HIDDEN)) + + +def test_bad_weight_shape_raises(): + op = Qwen3NextRMSNormOp() + with pytest.raises(ValueError, match="weight must be 1-D"): + op.forward_fp32(_rand((2, _HIDDEN), seed=28), _rand((_HIDDEN - 1,), seed=29)) + + +def test_gated_shape_mismatch_raises(): + op = Qwen3NextRMSNormGatedOp() + x, w = _rand((2, _HEAD_V_DIM), seed=30), _rand((_HEAD_V_DIM,), seed=31) + with pytest.raises(ValueError, match="gate must match x"): + op.forward_fp32(x, w, _rand((3, _HEAD_V_DIM), seed=32)) + + +# --------------------------------------------------------------------------- # +# 6. CUDA kernel: the offset is applied in fp32 inside the kernel +# --------------------------------------------------------------------------- # +_CUDA_RMSNORM = False +if torch.cuda.is_available(): # pragma: no branch - probe only + try: + from rl_engine.backends.extension import _C, _EXT_AVAILABLE + + _CUDA_RMSNORM = ( + _EXT_AVAILABLE + and getattr(_C, "rmsnorm_api_version", None) == 2 + and all(hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx")) + ) + except ImportError: # pragma: no cover + _CUDA_RMSNORM = False + +requires_cuda_rmsnorm = pytest.mark.skipif( + not _CUDA_RMSNORM, reason="CUDA RMSNorm extension is not available" +) + + +@requires_cuda_rmsnorm +def test_cuda_offset_equals_explicit_fp32_weight(): + """``weight_offset=1.0`` must equal passing an fp32 ``1 + w``, bitwise. + + This is the correctness proof for doing the offset inside the kernel: it is + the same arithmetic as the fp32 reference, not an approximation of it. + """ + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_cuda + + torch.manual_seed(0) + x = torch.randn(512, _HIDDEN, device="cuda", dtype=torch.bfloat16) + w = torch.randn(_HIDDEN, device="cuda", dtype=torch.bfloat16) + + offset = rmsnorm_cuda(x, w, eps=_EPS, weight_offset=1.0) + explicit = rmsnorm_cuda(x, (1.0 + w.float()), eps=_EPS) + assert torch.equal(offset, explicit) + + +@requires_cuda_rmsnorm +def test_cuda_offset_is_not_folded_through_bfloat16(): + """Pre-rounding ``1 + w`` to bf16 must give a different answer. + + If this ever becomes equal, the offset has stopped being applied in fp32 and + the zero-centred contract is silently broken. + """ + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_cuda + + torch.manual_seed(0) + x = torch.randn(512, _HIDDEN, device="cuda", dtype=torch.bfloat16) + w = torch.randn(_HIDDEN, device="cuda", dtype=torch.bfloat16) + + offset = rmsnorm_cuda(x, w, eps=_EPS, weight_offset=1.0) + folded = rmsnorm_cuda(x, (1.0 + w.float()).bfloat16(), eps=_EPS) + assert not torch.equal(offset, folded) + + +@requires_cuda_rmsnorm +def test_cuda_default_offset_preserves_plain_convention(): + """An unset offset must leave the existing kernel behaviour untouched.""" + from rl_engine.backends.cuda.norm.rmsnorm import RMSNormCudaOp, rmsnorm_cuda + + torch.manual_seed(0) + x = torch.randn(256, _HEAD_V_DIM, device="cuda", dtype=torch.bfloat16) + w = torch.randn(_HEAD_V_DIM, device="cuda", dtype=torch.bfloat16) + assert torch.equal(RMSNormCudaOp().forward(x, w, eps=_EPS), rmsnorm_cuda(x, w, eps=_EPS)) + + +@requires_cuda_rmsnorm +@pytest.mark.parametrize("batch", [1, 2, 8, 16, 32, 48, 64, 512]) +def test_cuda_zero_centred_batch_invariance(batch): + """L1 for the zero-centred CUDA op across the RFC #428 concurrency axis.""" + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp + + op = Qwen3NextRMSNormCudaOp() + torch.manual_seed(0) + x = torch.randn(512, _HIDDEN, device="cuda", dtype=torch.bfloat16) + w = torch.randn(_HIDDEN, device="cuda", dtype=torch.bfloat16) + assert torch.equal(op.forward(x[:batch], w, eps=_EPS), op.forward(x, w, eps=_EPS)[:batch]) + + +@requires_cuda_rmsnorm +def test_cuda_zero_centred_within_tolerance_of_golden(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp + + torch.manual_seed(0) + x = torch.randn(256, _HIDDEN, device="cuda", dtype=torch.bfloat16) + w = torch.randn(_HIDDEN, device="cuda", dtype=torch.bfloat16) + got = Qwen3NextRMSNormCudaOp().forward(x, w, eps=_EPS) + ref = Qwen3NextRMSNormOp().forward_fp32(x, w, eps=_EPS) + torch.testing.assert_close(got.float(), ref, **_forward_tol(torch.bfloat16)) + + +@requires_cuda_rmsnorm +def test_cuda_zero_centred_backward_is_offset_aware(): + """dx must see (1 + w); dw is offset-independent since d/dw (1+w) == d/dw w.""" + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_cuda + + torch.manual_seed(0) + x = torch.randn(128, _HEAD_V_DIM, device="cuda", dtype=torch.float32) + w = torch.randn(_HEAD_V_DIM, device="cuda", dtype=torch.float32) + dy = torch.randn(128, _HEAD_V_DIM, device="cuda", dtype=torch.float32) + + grads = {} + for name, off in (("plain", 0.0), ("zero_centred", 1.0)): + xg = x.clone().requires_grad_(True) + wg = w.clone().requires_grad_(True) + rmsnorm_cuda(xg, wg, eps=_EPS, weight_offset=off).backward(dy.clone()) + grads[name] = (xg.grad.clone(), wg.grad.clone()) + + assert torch.isfinite(grads["zero_centred"][0]).all() + assert not torch.equal(grads["plain"][0], grads["zero_centred"][0]) # dx differs + assert torch.equal(grads["plain"][1], grads["zero_centred"][1]) # dw does not + + +# --------------------------------------------------------------------------- # +# 8. `__call__` is the documented entry point and must agree with `forward` +# --------------------------------------------------------------------------- # +def test_call_matches_forward(): + x, w = _rand((4, _HIDDEN), seed=40), _rand((_HIDDEN,), seed=41) + op = Qwen3NextRMSNormOp() + assert torch.equal(op(x, w, eps=_EPS), op.forward(x, w, eps=_EPS)) + + +@pytest.mark.parametrize("cls", [Qwen3NextRMSNormGatedOp, Qwen3NextRMSNormGatedHFOp]) +def test_gated_call_matches_forward(cls): + x = _rand((4, _HEAD_V_DIM), seed=42) + w = _rand((_HEAD_V_DIM,), seed=43) + gate = _rand((4, _HEAD_V_DIM), seed=44) + op = cls() + assert torch.equal(op(x, w, gate, eps=_EPS), op.forward(x, w, gate, eps=_EPS)) + + +def test_zero_centred_op_inherits_the_plain_reference(): + """The only difference from the plain op is the weight convention. + + Pins the inheritance: if the subclass ever grows its own `_rms_norm`, this + stops being true and the two references can drift apart silently. + """ + from rl_engine.reference.norm.rms_norm import NativeRMSNormOp + + assert issubclass(Qwen3NextRMSNormOp, NativeRMSNormOp) + assert Qwen3NextRMSNormOp.weight_offset == 1.0 + assert NativeRMSNormOp.weight_offset == 0.0 + assert Qwen3NextRMSNormOp._rms_norm.__func__ is NativeRMSNormOp._rms_norm.__func__ + + +def test_plain_reference_is_unchanged_by_the_offset_plumbing(): + """Adding `weight_offset` to the base must not perturb the plain path. + + `0.0 + w` rewrites -0.0 to +0.0, which torch.equal does not notice, so this + compares the raw bits. + """ + from rl_engine.reference.norm.rms_norm import NativeRMSNormOp + + x = _rand((2, _HEAD_V_DIM), seed=45) + w = torch.zeros(_HEAD_V_DIM) + w[0] = -0.0 + got = NativeRMSNormOp().forward_fp32(x, w, eps=_EPS) + x_f = x.float() + from rl_engine.reference.norm.rms_norm import shape_invariant_rstd + + expected = (x_f * shape_invariant_rstd(x_f, _EPS).unsqueeze(-1)) * w.float() + assert torch.equal(got.view(torch.int32), expected.view(torch.int32)) + + +@requires_cuda_rmsnorm +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +def test_cuda_plain_signed_zero_is_preserved(dtype): + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_cuda + + x = torch.ones(2, 128, device="cuda", dtype=dtype) + x[1].neg_() + weight = torch.full((128,), -0.0, device="cuda", dtype=dtype) + actual = rmsnorm_cuda(x, weight) + expected = x * weight + bits = torch.int32 if dtype == torch.float32 else torch.int16 + assert torch.equal(actual.view(bits), expected.view(bits)) + + +@requires_cuda_rmsnorm +@pytest.mark.parametrize("offset", [0.0, 1.0]) +@pytest.mark.parametrize("shape", [(512, 128), (2, 3, 2048)]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_cuda_parameter_contributions_reproduce_backward(offset, shape, dtype): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp, RMSNormCudaOp + from rl_engine.ops.autograd.vjp_fp32 import reduce_rows_fp32 + + torch.manual_seed(128) + x = torch.randn(shape, device="cuda", dtype=dtype) + weight = torch.randn(shape[-1], device="cuda", dtype=dtype, requires_grad=True) + upstream = torch.randn_like(x) + op = RMSNormCudaOp() if offset == 0.0 else Qwen3NextRMSNormCudaOp() + op(x, weight).backward(upstream) + rows = op.parameter_vjp_contributions_fp32(x=x, weight=weight, grad_output=upstream)["weight"] + folded = reduce_rows_fp32(rows.reshape(-1, shape[-1])).to(dtype) + assert torch.equal(weight.grad, folded) + + +def test_zero_centred_cuda_constructor_rejects_missing_extension(monkeypatch): + from rl_engine.backends.cuda.norm import rmsnorm + + monkeypatch.setattr(rmsnorm, "_EXT_AVAILABLE", False) + monkeypatch.setattr(rmsnorm, "_C", None) + with pytest.raises(RuntimeError, match="requires the compiled"): + rmsnorm.Qwen3NextRMSNormCudaOp() diff --git a/tests/models/qwen3_next/test_qwen3_next_norm_reuse_check.py b/tests/models/qwen3_next/test_qwen3_next_norm_reuse_check.py new file mode 100644 index 000000000..f77b86bf7 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_norm_reuse_check.py @@ -0,0 +1,49 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Keep providers with incompatible weight semantics out of norm evidence.""" + +import sys +import types + +import pytest +import torch + +from tools.validation.models import qwen3_next_norm_reuse_check as reuse + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("uses_effective_weight", [False, True]) +def test_megatron_zero_centered_admission(monkeypatch, dtype, uses_effective_weight): + """Reject the 52fbcbc forward behavior while admitting a corrected provider.""" + module = types.ModuleType("megatron.core.transformer.custom_layers.batch_invariant_kernels") + + class BatchInvariantRMSNormFn: + @staticmethod + def apply(x, weight, eps, zero_centered_gamma): + weight_eff = weight + 1.0 if zero_centered_gamma else weight + scale = weight_eff if uses_effective_weight else weight + normalized = x.float() * torch.rsqrt(x.float().square().mean(-1, keepdim=True) + eps) + return (normalized * scale.float()).to(x.dtype) + + module.BatchInvariantRMSNormFn = BatchInvariantRMSNormFn + monkeypatch.setitem(sys.modules, module.__name__, module) + monkeypatch.setattr(reuse, "DEV", "cpu") + monkeypatch.setattr(reuse, "DT", dtype) + factory = dict(reuse._c1_candidates(32))["megatron"] + monkeypatch.setattr(reuse, "_c1_candidates", lambda hidden: [("megatron", factory)]) + + candidate = reuse.candidates("qwen3_next_rms_norm")["megatron"] + if not uses_effective_weight: + assert "failed the zero-weight probe" in candidate["unavailable"] + assert "fn" not in candidate + return + + assert candidate["backward"] == "full" + x = torch.ones(2, 32, dtype=dtype, requires_grad=True) + weight = torch.zeros(32, dtype=dtype, requires_grad=True) + output = candidate["fn"](x, weight) + assert torch.all(output != 0) + output.sum().backward() + assert x.grad is not None + assert weight.grad is not None diff --git a/tests/ops/norm/test_rms_norm.py b/tests/ops/norm/test_rms_norm.py index 90d7bab23..352612805 100644 --- a/tests/ops/norm/test_rms_norm.py +++ b/tests/ops/norm/test_rms_norm.py @@ -17,10 +17,11 @@ try: from rl_engine.backends.extension import _C, _EXT_AVAILABLE - # The same two symbols RMSNormCudaOp.__init__ requires; a build that has them - # dispatches to the CUDA op, so the dispatch test must agree with that guard. - _HAS_CUDA_RMSNORM = _EXT_AVAILABLE and all( - hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx") + # Keep dispatch expectations aligned with the constructor's capability guard. + _HAS_CUDA_RMSNORM = ( + _EXT_AVAILABLE + and getattr(_C, "rmsnorm_api_version", None) == 2 + and all(hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx")) ) except ImportError: # pragma: no cover - import can fail when the extension is not built. _HAS_CUDA_RMSNORM = False @@ -251,7 +252,7 @@ def test_backward_batch_invariance_slice(): # 9b. The CUDA backend must report itself unavailable by failing construction, # which is the seam the registry uses to fall back (see _get_or_create_backend). def test_cuda_op_construction_fails_without_extension(monkeypatch): - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) monkeypatch.setattr(cuda_rmsnorm, "_C", None) @@ -260,10 +261,10 @@ def test_cuda_op_construction_fails_without_extension(monkeypatch): def test_cuda_op_construction_fails_when_symbols_missing(monkeypatch): - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm class _WithoutRMSNorm: # a built extension that lacks the rmsnorm symbols - pass + rmsnorm_api_version = 2 monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr(cuda_rmsnorm, "_C", _WithoutRMSNorm()) @@ -273,8 +274,8 @@ class _WithoutRMSNorm: # a built extension that lacks the rmsnorm symbols def test_registry_falls_back_to_native_without_extension(monkeypatch): """A CUDA-first priority list must still resolve on a build without _C.""" - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) monkeypatch.setattr(cuda_rmsnorm, "_C", None) @@ -286,30 +287,66 @@ def test_registry_falls_back_to_native_without_extension(monkeypatch): def test_registry_falls_back_when_required_symbol_is_missing(monkeypatch, missing): from types import SimpleNamespace - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry symbols = {name: object() for name in ("rmsnorm_forward", "rmsnorm_backward_dx")} del symbols[missing] + symbols["rmsnorm_api_version"] = 2 monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr(cuda_rmsnorm, "_C", SimpleNamespace(**symbols)) assert isinstance(KernelRegistry().get_op("rms_norm", device="cuda"), NativeRMSNormOp) -def test_registry_cuda_requires_only_used_symbols_and_cpu_stays_native(monkeypatch): +@pytest.mark.parametrize("api_version", [None, 1, 3]) +def test_cuda_rejects_incompatible_rmsnorm_api_before_dispatch(monkeypatch, api_version): + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry + + class _LegacyRMSNorm: + def rmsnorm_forward(self, x, weight, eps): + pytest.fail("an incompatible extension must not be invoked") + + def rmsnorm_backward_dx(self, dy, x, weight, rstd): + pytest.fail("an incompatible extension must not be invoked") + + extension = _LegacyRMSNorm() + if api_version is not None: + extension.rmsnorm_api_version = api_version + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) + monkeypatch.setattr(cuda_rmsnorm, "_C", extension) + + for op_type in (cuda_rmsnorm.RMSNormCudaOp, cuda_rmsnorm.Qwen3NextRMSNormCudaOp): + with pytest.raises(RuntimeError, match="RMSNorm API version 2.*Rebuild"): + op_type() + + registry_op = KernelRegistry().get_op("rms_norm", device="cuda") + assert isinstance(registry_op, NativeRMSNormOp) + x, weight = torch.ones(2, 8), torch.ones(8) + assert torch.isfinite(registry_op(x, weight)).all() + with pytest.raises(RuntimeError, match="RMSNorm API version 2.*Rebuild"): + rmsnorm_cuda(x, weight, weight_offset=1.0) + + +def test_registry_cuda_requires_current_api_and_only_used_symbols(monkeypatch): from types import SimpleNamespace - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr( cuda_rmsnorm, "_C", - SimpleNamespace(rmsnorm_forward=object(), rmsnorm_backward_dx=object()), + SimpleNamespace( + rmsnorm_api_version=2, + rmsnorm_forward=object(), + rmsnorm_backward_dx=object(), + ), ) registry = KernelRegistry() assert isinstance(registry.get_op("rms_norm", device="cuda"), RMSNormCudaOp) + cuda_rmsnorm.Qwen3NextRMSNormCudaOp() assert isinstance(registry.get_op("rms_norm", device="cpu"), NativeRMSNormOp) diff --git a/tools/validation/models/plot_qwen3_next_norm_evidence.py b/tools/validation/models/plot_qwen3_next_norm_evidence.py new file mode 100644 index 000000000..5067d4d7c --- /dev/null +++ b/tools/validation/models/plot_qwen3_next_norm_evidence.py @@ -0,0 +1,121 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Plot a report from tools/validation/models/qwen3_next_norm_evidence.py: one 2x2 figure per op. + + python tools/validation/models/plot_qwen3_next_norm_evidence.py report.json + +Writes figure[-].png beside the report. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 + +COLORS = ["#2a6fdb", "#7a7a7a", "#e07b39", "#3aa676", "#9b59b6", "#c0392b"] + + +def _short(name: str) -> str: + return name.replace(" (forward only)", "*").replace(" (cast-first)", "\n(cast-first)") + + +def plot_op(op: str, data: dict, title: str, out: Path) -> None: + names = list(data["row_invariance"]) + color = {n: COLORS[i % len(COLORS)] for i, n in enumerate(names)} + fig, axes = plt.subplots(2, 2, figsize=(14, 10), layout="constrained") + fig.suptitle(title, fontsize=13) + + for ax, key, label in ( + (axes[0, 0], "forward_us", "forward latency"), + (axes[0, 1], "backward_us", "backward latency (training-capable only)"), + ): + for name in names: + rows = sorted(int(r) for r in data["latency"][name]) + ys = [data["latency"][name][str(r)].get(key) for r in rows] + if all(y is None for y in ys): + continue + ax.plot(rows, ys, "o-", color=color[name], label=_short(name)) + ax.set_xscale("log", base=2) + ax.set_yscale("log") + ax.set_xlabel("rows") + ax.set_ylabel("µs (median)") + ax.set_title(label) + ax.grid(True, which="both", alpha=0.3) + ax.legend(fontsize=8) + + ax = axes[1, 0] + acc = data["accuracy"][max(data["accuracy"], key=int)] + metrics = [ + ("forward_max_abs", "forward\nmax |err|"), + ("dx_max_abs_over_absmax", "dx\nmax |err| / max"), + ("dweight_max_abs_over_absmax", "dweight\nmax |err| / max"), + ("dgate_max_abs_over_absmax", "dgate\nmax |err| / max"), + ] + metrics = [m for m in metrics if any(acc[n].get(m[0]) is not None for n in names)] + width = 0.8 / len(names) + for i, name in enumerate(names): + vals = [acc[name].get(m) for m, _ in metrics] + xs = [j + (i - (len(names) - 1) / 2) * width for j in range(len(metrics))] + ax.bar( + [x for x, v in zip(xs, vals) if v is not None], + [v for v in vals if v is not None], + width, + color=color[name], + label=_short(name), + ) + ax.set_xticks(range(len(metrics)), [label for _, label in metrics], fontsize=9) + ax.set_yscale("log") + ax.set_title(f"error vs FP64 golden ({max(data['accuracy'], key=int)} rows, BF16)") + ax.grid(True, axis="y", alpha=0.3) + ax.legend(fontsize=8) + + ax = axes[1, 1] + bi = data["row_invariance"] + checked = next(iter(bi.values()))["rows_checked"] + ys = range(len(names)) + fwd = [bi[n]["forward_rows_differing"] for n in names] + dx = [ + bi[n]["dx_rows_differing"] if bi[n]["dx_rows_differing"] is not None else 0 for n in names + ] + ax.barh([y - 0.2 for y in ys], fwd, 0.4, color="#2a6fdb", label="forward") + ax.barh([y + 0.2 for y in ys], dx, 0.4, color="#e07b39", label="dx") + for y, n, f, d in zip(ys, names, fwd, dx): + no_bwd = bi[n]["dx_rows_differing"] is None + ax.text(max(f, d) + 0.2, y, f"{f} / {'n/a' if no_bwd else d}", va="center", fontsize=8) + ax.set_yticks(list(ys), [_short(n) for n in names], fontsize=8) + ax.set_xlim(0, max(3, max(fwd + dx) * 1.4)) + ax.invert_yaxis() + ax.set_xlabel(f"rows differing (of {checked}; row alone vs inside a batch, bitwise)") + ax.set_title("row invariance (0 = batch-invariant)") + ax.legend(fontsize=8) + ax.grid(True, axis="x", alpha=0.3) + + fig.savefig(out, dpi=130) + print(f"wrote {out}") + + +def main() -> None: + path = Path(sys.argv[1]) + report = json.loads(path.read_text()) + env = report["environment"] + ops = report["ops"] + for op, data in ops.items(): + name = "figure.png" if len(ops) == 1 else f"figure-{op}.png" + hidden = data.get("hidden", report.get("hidden")) + title = ( + f"Qwen3-Next {op.replace('_', ' ')} — {env['gpu']}, hidden {hidden}, " + f"BF16, commit {report['git_commit'][:7]} (* forward only)" + ) + plot_op(op, data, title, path.parent / name) + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/qwen3_next_norm_evidence.py b/tools/validation/models/qwen3_next_norm_evidence.py new file mode 100644 index 000000000..34018c013 --- /dev/null +++ b/tools/validation/models/qwen3_next_norm_evidence.py @@ -0,0 +1,276 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Accuracy, row invariance and latency of the Qwen3-Next norms vs existing implementations. + +Writes one JSON report (RFC #428, reuse rule of #420): every candidate is compared with +an FP64 golden, checked for row invariance (a row computed alone vs inside a batch, +bitwise), and timed. Optional providers (transformers, vLLM, FlashInfer) are skipped +when they are not installed; the report says which ran. + + python tools/validation/models/qwen3_next_norm_evidence.py --out report.json + python tools/validation/models/plot_qwen3_next_norm_evidence.py report.json +""" + +from __future__ import annotations + +import argparse +import json +import platform +import statistics +import subprocess +import sys +from pathlib import Path +from typing import Any, Callable + +import torch + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp # noqa: E402 +from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormOp # noqa: E402 + +EPS = 1e-6 +HIDDEN = 2048 # Qwen3-Next hidden size (decoder and final norms) +TIMED_ROWS = (1024, 4096, 16384, 65536) + + +# --------------------------------------------------------------------------- # +# Candidates: name -> (forward fn(x, w) -> y, has_backward) +# --------------------------------------------------------------------------- # + + +def zero_centred_candidates() -> dict[str, dict[str, Any]]: + cands: dict[str, dict[str, Any]] = { + "rl-kernel CUDA": {"fn": Qwen3NextRMSNormCudaOp(), "backward": True}, + "rl-kernel PyTorch reference": {"fn": Qwen3NextRMSNormOp(), "backward": True}, + } + try: + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNorm + + module = Qwen3NextRMSNorm(HIDDEN, eps=EPS).cuda() + + def hf(x, w, module=module): + return torch.func.functional_call(module, {"weight": w}, (x,)) + + cands["transformers Qwen3NextRMSNorm"] = {"fn": hf, "backward": True} + except ImportError: + pass + try: + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + from vllm.model_executor.layers.layernorm import GemmaRMSNorm + + gemma = GemmaRMSNorm(HIDDEN, eps=EPS).cuda() + + def vllm_fwd(x, w, gemma=gemma): + gemma.weight.data = w.detach() + return gemma.forward_cuda(x) + + cands["vLLM GemmaRMSNorm (forward only)"] = {"fn": vllm_fwd, "backward": False} + except ImportError: + pass + try: + import flashinfer + + cands["FlashInfer gemma_rmsnorm (forward only)"] = { + "fn": lambda x, w: flashinfer.norm.gemma_rmsnorm(x, w, EPS), + "backward": False, + } + except ImportError: + pass + return cands + + +def zero_centred_golden(x: torch.Tensor, w: torch.Tensor) -> torch.Tensor: + x64 = x.double() + return x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + EPS) * (1.0 + w.double()) + + +# --------------------------------------------------------------------------- # +# Measurements +# --------------------------------------------------------------------------- # + + +def _inputs(rows: int, seed: int, dtype=torch.bfloat16): + g = torch.Generator(device="cuda").manual_seed(seed) + x = (torch.randn(rows, HIDDEN, device="cuda", generator=g) * 2).to(dtype) + w = (torch.randn(HIDDEN, device="cuda", generator=g) * 0.1).to(dtype) + up = torch.randn(rows, HIDDEN, device="cuda", generator=g).to(dtype) + return x, w, up + + +def _grads(fn, x, w, up, dtype=None): + xl = (x if dtype is None else x.to(dtype)).detach().clone().requires_grad_(True) + wl = (w if dtype is None else w.to(dtype)).detach().clone().requires_grad_(True) + out = fn(xl, wl) + out.backward(up if dtype is None else up.to(dtype)) + return out.detach(), xl.grad, wl.grad + + +def accuracy(cands, golden, rows: int, seed: int) -> dict[str, Any]: + x, w, up = _inputs(rows, seed) + ref_out, ref_dx, ref_dw = _grads(golden, x, w, up, torch.float64) + result = {} + for name, c in cands.items(): + entry: dict[str, Any] = {} + if c["backward"]: + out, dx, dw = _grads(c["fn"], x, w, up) + for key, got, ref in (("dx", dx, ref_dx), ("dweight", dw, ref_dw)): + err = (got.double() - ref).abs().max().item() + entry[f"{key}_max_abs_over_absmax"] = err / ref.abs().max().item() + else: + with torch.no_grad(): + out = c["fn"](x, w) + err = (out.double() - ref_out).abs() + entry["forward_max_abs"] = err.max().item() + entry["forward_correctly_rounded_fraction"] = ( + (out == ref_out.to(out.dtype)).float().mean().item() + ) + result[name] = entry + return result + + +def row_invariance(cands, seeds=(3, 4, 5), rows: int = 4096, step: int = 16) -> dict[str, Any]: + """256 rows computed alone vs the same rows inside a batch, bitwise.""" + + result = {} + for name, c in cands.items(): + fwd_bad = dx_bad = checked = 0 + for seed in seeds: + x, w, up = _inputs(rows, seed) + if c["backward"]: + full_out, full_dx, _ = _grads(c["fn"], x, w, up) + else: + with torch.no_grad(): + full_out = c["fn"](x, w) + for i in range(0, rows, step): + sl = slice(i, i + 1) + if c["backward"]: + out, dx, _ = _grads(c["fn"], x[sl], w, up[sl]) + dx_bad += not torch.equal(dx[0], full_dx[i]) + else: + with torch.no_grad(): + out = c["fn"](x[sl], w) + fwd_bad += not torch.equal(out[0], full_out[i]) + checked += 1 + result[name] = { + "rows_checked": checked, + "forward_rows_differing": fwd_bad, + "dx_rows_differing": dx_bad if c["backward"] else None, + "batch_rows": rows, + "seeds": list(seeds), + } + return result + + +def _time_us(fn: Callable[[], Any], warmup: int = 10, iters: int = 50) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(iters): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1e3) + return statistics.median(samples) + + +def latency(cands, rows_list=TIMED_ROWS) -> dict[str, Any]: + result: dict[str, Any] = {name: {} for name in cands} + for rows in rows_list: + x, w, up = _inputs(rows, seed=11) + for name, c in cands.items(): + row: dict[str, float] = {} + with torch.no_grad(): + row["forward_us"] = _time_us(lambda f=c["fn"]: f(x, w)) + if c["backward"]: + xl = x.detach().clone().requires_grad_(True) + wl = w.detach().clone().requires_grad_(True) + out = c["fn"](xl, wl) + row["backward_us"] = _time_us( + lambda o=out: torch.autograd.grad(o, (xl, wl), up, retain_graph=True) + ) + result[name][str(rows)] = row + return result + + +# --------------------------------------------------------------------------- # + + +def _git(*args: str) -> str: + try: + return subprocess.check_output(["git", *args], cwd=REPO_ROOT, text=True).strip() + except (OSError, subprocess.CalledProcessError): + return "" + + +def environment() -> dict[str, Any]: + env = { + "gpu": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "python": platform.python_version(), + } + for mod in ("transformers", "vllm", "flashinfer"): + try: + env[mod] = __import__(mod).__version__ + except ImportError: + env[mod] = None + return env + + +def run_op(name: str, cands, golden) -> dict[str, Any]: + print(f"[{name}] candidates: {', '.join(cands)}", flush=True) + report = { + "accuracy": {str(r): accuracy(cands, golden, r, seed=r) for r in (257, 4096)}, + "row_invariance": row_invariance(cands), + "latency": latency(cands), + } + for cand, entry in report["row_invariance"].items(): + print(f" BI {cand}: {entry}", flush=True) + return report + + +def build_report(ops: dict[str, tuple[Callable, Callable]]) -> dict[str, Any]: + torch.backends.cuda.matmul.allow_tf32 = False + report = { + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": _git("rev-parse", "HEAD") or "unknown", + "git_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": environment(), + "hidden": HIDDEN, + "eps": EPS, + "dtype": "bfloat16", + "ops": {}, + } + for name, (make_cands, golden) in ops.items(): + report["ops"][name] = run_op(name, make_cands(), golden) + return report + + +OPS = {"zero_centred_rmsnorm": (zero_centred_candidates, zero_centred_golden)} + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True) + args = parser.parse_args() + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + report = build_report(OPS) + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + print(f"wrote {args.out} (commit {report['git_commit'][:7]}, dirty={report['git_dirty']})") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/qwen3_next_norm_reuse_check.py b/tools/validation/models/qwen3_next_norm_reuse_check.py new file mode 100644 index 000000000..5fd03b53c --- /dev/null +++ b/tools/validation/models/qwen3_next_norm_reuse_check.py @@ -0,0 +1,710 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Compare a Qwen3-Next norm with existing implementations: batch invariance, accuracy, gates. + +Checks, for the rl-kernel CUDA op and every existing implementation that imports here +(transformers, vLLM, FlashInfer, Liger, FLA, Transformer Engine, Megatron-LM from a source +tree): + +* ``bi``: bitwise batch invariance of the forward and the row gradients, and repeatability of + ``dweight``: every row computed alone vs inside full batches of three sizes (seeds 3-5); + the full workload-size batch vs sub-batches covering every row; eight probe rows at the + front, middle and back of batches of every size 1..9 and 2^k - 1, 2^k, 2^k + 1. Sparse probes + miss row-specific and large-batch dependence, which is why every row is checked. +* ``perf``: accuracy against FP64 (forward: fraction equal to the correctly rounded FP64 result; + gradients: max|err| / max|ref|) and median CUDA-event latency at the workload size. +* ``gates``: this repository's own C3/C4 gates + (``tools/validation/operators/check_forward_invariance.py`` and + ``tools/validation/operators/check_gradient_invariance.py``) with the norm manifest and the + arguments ``ci/scripts/run_ws1_gtest.sh`` uses), unchanged, with the CUDA candidate replaced by a + subclass of the rl-kernel op whose forward and backward call the other library. The subclass + keeps the op's ``parameter_vjp_contributions_fp32``, so the singleton-aggregate check + compares that library's ``dweight`` with the same FP32 row contributions. + + python tools/validation/models/qwen3_next_norm_reuse_check.py --op qwen3_next_rms_norm \\ + --out docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json \\ + [--checks bi,perf,gates] [--megatron-src /path/to/Megatron-LM] [--quick] + +Run on a clean tree on an otherwise idle GPU; the report records the commit, the tree state and +every library version. A figure is written next to the report when matplotlib is available. +""" + +from __future__ import annotations + +import argparse +import importlib.metadata as md +import json +import os +import statistics +import subprocess +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +import torch # noqa: E402 +import torch.nn.functional as F # noqa: E402 + +DEV, EPS, DT = "cuda", 1e-6, torch.bfloat16 +SHAPES = { + # op: (hidden, all-rows batch sizes, full batch, covering sub-batch sizes, perf rows) + "qwen3_next_rms_norm": (2048, (64, 257, 2048), 65536, (1, 7, 2048, 4097, 32768), 65536), + "rms_norm_gated": (128, (64, 257, 4097), 262144, (1, 7, 4097, 65536, 131072), 262144), +} +LIBS = ("torch", "triton", "transformers", "vllm", "flashinfer-python", "liger-kernel", "fla-core") +LIBS += ("transformer_engine", "megatron-core") +MANIFEST = "rl_engine/validation/models/qwen3_next_norm_manifest.json" + + +def _version(name: str) -> str | None: + try: + return md.version(name) + except md.PackageNotFoundError: + return None + + +# --------------------------------------------------------------------------- # +# Candidates: fn(x, w) for the zero-centred norm, fn(x, z, w) for the gated norm +# --------------------------------------------------------------------------- # + + +def _vllm_module(build): + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + return build() + + +def _c1_candidates(hidden: int) -> list[tuple[str, Any]]: + def ours(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp + + o = Qwen3NextRMSNormCudaOp() + return "rl-kernel Qwen3NextRMSNormCudaOp", lambda x, w: o(x, w, eps=EPS), "full" + + def reference(): + from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormOp + + o = Qwen3NextRMSNormOp() + return "rl-kernel PyTorch reference", lambda x, w: o(x, w, eps=EPS), "full" + + def transformers(): + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNorm + + m = Qwen3NextRMSNorm(hidden, eps=EPS).to(DEV, DT) + + def fn(x, w): + return torch.func.functional_call(m, {"weight": w}, (x,)) + + return f"transformers {_version('transformers')} Qwen3NextRMSNorm", fn, "full" + + def liger(): + from liger_kernel.ops.rms_norm import LigerRMSNormFunction + + def fn(x, w): + return LigerRMSNormFunction.apply(x, w, EPS, 1.0, "gemma", False) + + return f"Liger {_version('liger-kernel')} RMSNorm, offset 1, gemma", fn, "full" + + def fla(): + from fla.modules.layernorm import rms_norm + + def fn(x, w): + return rms_norm(x, (1.0 + w.float()).to(x.dtype), None, eps=EPS) + + return f"FLA {_version('fla-core')} rms_norm, weight passed as 1 + w", fn, "full" + + def te(): + import transformer_engine.pytorch as te_ + + m = te_.RMSNorm(hidden, eps=EPS, zero_centered_gamma=True, params_dtype=DT, device=DEV) + + def fn(x, w): + # TE's backward reads its own parameter, so copy w in (functional_call gives + # a wrong dx); dweight is then TE's parameter gradient and is not checked. + with torch.no_grad(): + m.weight.copy_(w) + return m(x) + + return f"TE {_version('transformer_engine')} RMSNorm(zero_centered_gamma)", fn, "x-only" + + def megatron(): + from megatron.core.transformer.custom_layers.batch_invariant_kernels import ( + BatchInvariantRMSNormFn, + ) + + def fn(x, w): + return BatchInvariantRMSNormFn.apply(x, w, EPS, True) + + # Some revisions accept the flag but still multiply by the uncentered weight. + with torch.no_grad(): + probe_x = torch.ones(2, hidden, device=DEV, dtype=DT) + probe_w = torch.zeros(hidden, device=DEV, dtype=DT) + expected = (probe_x.float() * (1.0 + EPS) ** -0.5).to(DT) + if not torch.allclose(fn(probe_x, probe_w), expected): + raise RuntimeError( + "Megatron zero_centered_gamma=True failed the zero-weight probe; " + "excluded from the zero-centered comparison" + ) + + return "Megatron BatchInvariantRMSNormFn(zero_centered_gamma=True)", fn, "full" + + def flashinfer(): + import flashinfer as fi + + def fn(x, w): + return fi.gemma_rmsnorm(x, w, EPS) + + return f"FlashInfer {_version('flashinfer-python')} gemma_rmsnorm", fn, "none" + + def vllm(): + from vllm.model_executor.layers.layernorm import GemmaRMSNorm + + m = _vllm_module(lambda: GemmaRMSNorm(hidden, eps=EPS)).to(DEV, DT) + + def fn(x, w): + m.weight.data.copy_(w) + return m.forward_cuda(x) + + return f"vLLM {_version('vllm')} GemmaRMSNorm.forward_cuda", fn, "none" + + return [ + ("rl_kernel", ours), + ("reference", reference), + ("transformers", transformers), + ("liger", liger), + ("fla", fla), + ("transformer_engine", te), + ("megatron", megatron), + ("flashinfer", flashinfer), + ("vllm", vllm), + ] + + +def _gated_candidates(hidden: int) -> list[tuple[str, Any]]: + def ours(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + o = Qwen3NextRMSNormGatedCudaOp() + return "rl-kernel Qwen3NextRMSNormGatedCudaOp", lambda x, z, w: o(x, w, z, eps=EPS), "full" + + def reference(): + from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormGatedOp + + o = Qwen3NextRMSNormGatedOp() + return "rl-kernel PyTorch reference", lambda x, z, w: o(x, w, z, eps=EPS), "full" + + def transformers(): + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNormGated + + m = Qwen3NextRMSNormGated(hidden, eps=EPS).to(DEV, DT) + + def fn(x, z, w): + return torch.func.functional_call(m, {"weight": w}, (x, z)) + + return f"transformers {_version('transformers')} Qwen3NextRMSNormGated", fn, "full" + + def fla_lg(): + from fla.modules.layernorm_gated import rmsnorm_fn + + def fn(x, z, w): + return rmsnorm_fn(x, w, None, z=z, eps=EPS, group_size=None, norm_before_gate=True) + + return f"FLA {_version('fla-core')} layernorm_gated.rmsnorm_fn", fn, "full" + + def fla_fused(): + from fla.modules.fused_norm_gate import rms_norm_gated + + def fn(x, z, w): + return rms_norm_gated(x, z, w, None, activation="silu", eps=EPS) + + return f"FLA {_version('fla-core')} fused_norm_gate.rms_norm_gated", fn, "full" + + def vllm(): + from vllm.model_executor.layers.layernorm import RMSNormGated + + def build(): + return RMSNormGated( + hidden, eps=EPS, group_size=None, norm_before_gate=True, activation="silu" + ) + + m = _vllm_module(build).to(DEV, DT) + + def fn(x, z, w): + m.weight.data.copy_(w) + return m.forward_cuda(x, z) + + return f"vLLM {_version('vllm')} RMSNormGated.forward_cuda", fn, "none" + + return [ + ("rl_kernel", ours), + ("reference", reference), + ("transformers", transformers), + ("fla_layernorm_gated", fla_lg), + ("fla_fused_norm_gate", fla_fused), + ("vllm", vllm), + ] + + +def candidates(op: str) -> dict[str, dict[str, Any]]: + hidden = SHAPES[op][0] + found = _c1_candidates(hidden) if op == "qwen3_next_rms_norm" else _gated_candidates(hidden) + out: dict[str, dict[str, Any]] = {} + for name, factory in found: + try: + source, fn, backward = factory() + out[name] = {"source": source, "fn": fn, "backward": backward} + except Exception as exc: # noqa: BLE001 - optional library missing or unusable + out[name] = {"unavailable": f"{type(exc).__name__}: {exc}"[:300]} + return out + + +def _inputs(op: str, seed: int, n: int) -> dict[str, torch.Tensor]: + hidden = SHAPES[op][0] + g = torch.Generator(device=DEV).manual_seed(seed) + x = (torch.randn(n, hidden, device=DEV, generator=g) * 2).to(DT) + if op == "qwen3_next_rms_norm": + return {"x": x, "w": (torch.randn(hidden, device=DEV, generator=g) * 0.1).to(DT)} + z = torch.randn(n, hidden, device=DEV, generator=g).to(DT) + return {"x": x, "z": z, "w": (1 + torch.randn(hidden, device=DEV, generator=g) * 0.1).to(DT)} + + +def _reference(op: str, inputs: dict[str, torch.Tensor]): + x = inputs["x"].double() + n = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) + if op == "qwen3_next_rms_norm": + return n * (1 + inputs["w"].double()) + return n * inputs["w"].double() * F.silu(inputs["z"].double()) + + +# --------------------------------------------------------------------------- # +# Batch invariance +# --------------------------------------------------------------------------- # + + +def _run(fn, inputs, rows, grads, ups, idx=None): + args = {} + for k, v in inputs.items(): + if idx is not None and k in rows: + v = v.index_select(0, idx) + args[k] = v.detach().clone().requires_grad_(True) if k in grads else v + out = fn(**args) + if grads: + out.backward(ups if idx is None else ups.index_select(0, idx)) + return out.detach(), {k: args[k].grad for k in grads} + + +def _sweep_sizes(n: int) -> list[int]: + sizes = set(range(1, min(n, 9) + 1)) + k = 16 + while k <= n: + sizes.update(v for v in (k - 1, k, k + 1) if v <= n) + k *= 2 + return sorted(sizes | {n}) + + +def batch_invariance(op: str, cand: dict[str, Any], quick: bool) -> dict[str, Any]: + _, full_sizes, big, subs, _ = SHAPES[op] + if quick: + full_sizes, big, subs = (17, 64), 512, (1, 7, 64) + fn = cand["fn"] + rows = ["x"] if op == "qwen3_next_rms_norm" else ["x", "z"] + grads = {"full": rows + ["w"], "x-only": rows, "none": []}[cand["backward"]] + row_grads = [k for k in rows if k in grads] + + def ups_for(seed, n): + g = torch.Generator(device=DEV).manual_seed(seed + 100) + return torch.randn(n, SHAPES[op][0], device=DEV, generator=g).to(DT) + + res: dict[str, Any] = {"all_rows": {}, "dweight_repeatable": None} + ok = True + for n in full_sizes: + fwd = grd = 0 + repeat = True + for seed in (3, 4, 5): + inp, ups = _inputs(op, seed, n), ups_for(seed, n) + fo, fg = _run(fn, inp, rows, grads, ups) + for i in range(n): + o, g = _run(fn, inp, rows, grads, ups, torch.tensor([i], device=DEV)) + fwd += not torch.equal(o[0], fo[i]) + grd += not all(torch.equal(g[k][0], fg[k][i]) for k in row_grads) + if "w" in grads: + _, again = _run(fn, inp, rows, grads, ups) + repeat &= torch.equal(again["w"], fg["w"]) + if "w" in grads: + res["dweight_repeatable"] = repeat and res["dweight_repeatable"] is not False + res["all_rows"][str(n)] = {"rows": 3 * n, "fwd_differ": fwd, "rowgrad_differ": grd} + ok &= fwd == 0 and grd == 0 and repeat + + fwd = grd = checked = 0 + for seed in (3, 4): + inp, ups = _inputs(op, seed, big), ups_for(seed, big) + fo, fg = _run(fn, inp, rows, grads, ups) + for sb in subs: + step = sb if sb >= 1024 else max(sb, big // 512) + for start in range(0, big, step): + idx = torch.arange(start, min(start + sb, big), device=DEV) + o, g = _run(fn, inp, rows, grads, ups, idx) + checked += 1 + fwd += not torch.equal(o, fo.index_select(0, idx)) + grd += not all(torch.equal(g[k], fg[k].index_select(0, idx)) for k in row_grads) + del inp, ups, fo, fg + torch.cuda.empty_cache() + res["full_vs_sub_batches"] = { + "full_batch": big, + "sub_batch_sizes": list(subs), + "sub_batches": checked, + "fwd_differ": fwd, + "rowgrad_differ": grd, + } + ok &= fwd == 0 and grd == 0 + + n = full_sizes[-1] + sizes = _sweep_sizes(n) + probes = sorted({0, n - 1, *list(range(n // 8 + 3, n - 1, n // 8))[:6]}) + compared = failed = 0 + for seed in (3, 4, 5): + inp, ups = _inputs(op, seed, n), ups_for(seed, n) + alone = {r: _run(fn, inp, rows, grads, ups, torch.tensor([r], device=DEV)) for r in probes} + gen = torch.Generator().manual_seed(seed) + for m in sizes: + for r in probes: + others = torch.randperm(n, generator=gen) + others = others[others != r][: m - 1] + for pos in sorted({0, (m - 1) // 2, m - 1}): + idx = torch.cat([others[:pos], torch.tensor([r]), others[pos:]]).to(DEV) + o, g = _run(fn, inp, rows, grads, ups, idx) + good = torch.equal(o[pos], alone[r][0][0]) + good &= all(torch.equal(g[k][pos], alone[r][1][k][0]) for k in row_grads) + compared += 1 + failed += not good + res["batch_size_sweep"] = { + "batch_sizes": sizes, + "probe_rows": probes, + "comparisons": compared, + "differ": failed, + } + ok &= failed == 0 + res["batch_invariant"] = bool(ok) + return res + + +# --------------------------------------------------------------------------- # +# Accuracy and latency +# --------------------------------------------------------------------------- # + + +def _time_us(fn, warmup=10, iters=50) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(iters): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1e3) + return statistics.median(samples) + + +def accuracy_latency(op: str, cand: dict[str, Any], quick: bool) -> dict[str, Any]: + n = 4096 if quick else SHAPES[op][4] + inp = _inputs(op, 11, n) + keys = list(inp) + gk = {"full": keys, "x-only": [k for k in keys if k != "w"], "none": []}[cand["backward"]] + fn = cand["fn"] + leaves64 = {k: v.double().requires_grad_(k in gk) for k, v in inp.items()} + ref = _reference(op, leaves64) + with torch.no_grad(): + y = fn(**inp) + out: dict[str, Any] = { + "rows": n, + "fwd_correctly_rounded": (y == ref.detach().to(DT)).float().mean().item(), + "fwd_max_abs_err": (y.double() - ref.detach()).abs().max().item(), + "fwd_us": _time_us(lambda: fn(**inp)), + } + if gk: + gen = torch.Generator(device=DEV).manual_seed(5) + dy = torch.randn(n, SHAPES[op][0], device=DEV, generator=gen).to(DT) + rg = torch.autograd.grad(ref, [leaves64[k] for k in gk], dy.double()) + lv = {k: v.detach().clone().requires_grad_(k in gk) for k, v in inp.items()} + gs = torch.autograd.grad(fn(**lv), [lv[k] for k in gk], dy) + out["grad_rel_err"] = { + k: ((g.double() - r).abs().max() / r.abs().max()).item() for k, g, r in zip(gk, gs, rg) + } + + def step(): + lv = {k: v.detach().clone().requires_grad_(k in gk) for k, v in inp.items()} + torch.autograd.grad(fn(**lv), [lv[k] for k in gk], dy) + + out["fwd_bwd_us"] = _time_us(step) + return out + + +# --------------------------------------------------------------------------- # +# The repository's C3/C4 gates with an existing implementation swapped in +# --------------------------------------------------------------------------- # + + +def _gate_child(which: str, op: str, impl: str) -> None: + """Run the canonical C3/C4 CLI with the CUDA candidate replaced by ``impl``.""" + + import importlib.util + + from rl_engine.backends.cuda.norm import rmsnorm as R + + path = REPO_ROOT / "tools/validation/operators" / f"check_{which}_invariance.py" + spec = importlib.util.spec_from_file_location("gate", path) + gate = importlib.util.module_from_spec(spec) + sys.argv = [ + str(path), + "--manifest", + MANIFEST, + "--op", + op, + "--candidate", + "cuda", + "--backend-profile", + "cuda_bf16", + "--hidden", + "2048", + "--head-dim", + "128", + ] + spec.loader.exec_module(gate) + if impl != "rl_kernel": + fn = candidates(op)[impl]["fn"] + if op == "qwen3_next_rms_norm": + + class Swapped(R.Qwen3NextRMSNormCudaOp): + def forward(self, x, weight, *, eps=EPS): + return fn(x.reshape(-1, x.shape[-1]), weight).view_as(x) + + else: + + class Swapped(R.Qwen3NextRMSNormGatedCudaOp): + def __call__(self, x, weight, gate, *, eps=EPS): + return self.forward(x, weight, gate, eps=eps) + + def forward(self, x, weight, gate, *, eps=EPS): + return fn(x, gate, weight) + + original = gate.load_adapter_operator + + def load_adapter_operator(op_name, candidate): + return Swapped() if candidate == "cuda" else original(op_name, candidate) + + gate.load_adapter_operator = load_adapter_operator + gate.main() + + +def _clean(lines: list[str], megatron_src: str | None) -> list[str]: + """Drop warnings and replace machine-specific paths, so reports carry no local paths.""" + + subs = [(str(REPO_ROOT), ""), (sys.prefix, ""), (str(Path.home()), "")] + if megatron_src: + subs.insert(0, (str(Path(megatron_src).resolve()), "")) + out = [] + for line in lines: + if "Warning" in line or "warnings.warn" in line: + continue + for path, name in subs: + line = line.replace(path, name) + out.append(line) + return out + + +def contract_gates(op: str, impl: str, megatron_src: str | None) -> dict[str, Any]: + if not (REPO_ROOT / MANIFEST).exists(): + # The Qwen3-Next C3/C4 gate adapters and manifest arrive with the gated-norm PR. + return {"unavailable": f"{MANIFEST} is not on this branch"} + out: dict[str, Any] = {} + for which in ("forward", "gradient"): + cmd = [ + sys.executable, + __file__, + "--op", + op, + "--out", + os.devnull, + "--gate-child", + which, + impl, + ] + if megatron_src: + cmd += ["--megatron-src", megatron_src] + proc = subprocess.run( + cmd, capture_output=True, text=True, env={**os.environ, "RL_KERNEL_REQUIRE_EXT": "1"} + ) + lines = _clean(proc.stdout.splitlines(), megatron_src) + summary = next((ln for ln in lines if ln.startswith("op=")), "") + out[which] = { + "returncode": proc.returncode, + "passed": "passed=True" in summary and proc.returncode == 0, + "failed_lines": [ln.strip() for ln in lines if "passed=False" in ln][:10], + "singleton_aggregate": [ln.strip() for ln in lines if "singleton_aggregate pair" in ln], + "stderr_tail": ( + _clean(proc.stderr.strip().splitlines(), megatron_src)[-3:] + if proc.returncode + else [] + ), + } + return out + + +# --------------------------------------------------------------------------- # +# Report and figure +# --------------------------------------------------------------------------- # + + +def plot(report: dict[str, Any], path: Path) -> None: + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + except ImportError: + return + entries = [(n, e) for n, e in report["results"].items() if "unavailable" not in e] + names = [e["source"] for _, e in entries] + ys = list(range(len(entries))) + fig, axes = plt.subplots(1, 3, figsize=(22, 0.55 * len(entries) + 3), layout="constrained") + fig.suptitle( + f"{report['op']} vs existing implementations — {report['environment']['gpu']}, " + f"commit {report['rl_kernel_commit'][:7]}" + ) + ax = axes[0] + for key, off, label in (("fwd_us", -0.2, "forward"), ("fwd_bwd_us", 0.2, "fwd + bwd")): + vals = [e.get("accuracy_latency", {}).get(key, float("nan")) for _, e in entries] + ax.barh([y + off for y in ys], vals, 0.4, label=label) + ax.set_xscale("log") + ax.set_yticks(ys, names, fontsize=7) + ax.invert_yaxis() + ax.set_xlabel("µs, median") + ax.set_title("latency at the workload size") + ax.legend(fontsize=7) + ax.grid(True, axis="x", alpha=0.3) + ax = axes[1] + acc = [e.get("accuracy_latency", {}) for _, e in entries] + cr = [a.get("fwd_correctly_rounded", float("nan")) for a in acc] + ax.barh(ys, [max(1 - c, 1e-7) for c in cr]) + for y, c, a in zip(ys, cr, acc): + grads = a.get("grad_rel_err", {}) + worst = f"; worst grad err {max(grads.values()):.1e}" if grads else "" + ax.text(1.5e-7, y, f"{c:.4%} correctly rounded{worst}", va="center", fontsize=7) + ax.set_xscale("log") + ax.set_xlim(1e-7, 1) + ax.set_yticks(ys, [""] * len(entries)) + ax.invert_yaxis() + ax.set_title("forward: fraction NOT equal to the correctly rounded FP64 result") + ax = axes[2] + for y, (_, e) in zip(ys, entries): + bi = e.get("batch_invariance") + gates = e.get("contract_gates") + text = "batch-invariant: " + ( + "n/a" if bi is None else ("yes" if bi["batch_invariant"] else "NO") + ) + if gates and "unavailable" not in gates: + text += " | C3/C4 gates: " + ( + "pass" if all(g["passed"] for g in gates.values()) else "FAIL" + ) + good = bi is not None and bi["batch_invariant"] + ax.text(0.02, y, text, va="center", fontsize=8, color="#3aa676" if good else "#d14b4b") + ax.set_ylim(len(entries) - 0.5, -0.5) + ax.axis("off") + ax.set_title("every-row batch invariance and this repository's gates") + fig.savefig(path.with_suffix(".png"), dpi=120) + + +def _git(*args: str) -> str: + return subprocess.run( + ["git", *args], cwd=REPO_ROOT, capture_output=True, text=True + ).stdout.strip() + + +def main() -> None: + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument("--op", required=True, choices=sorted(SHAPES)) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--checks", default="bi,perf,gates") + parser.add_argument("--megatron-src", default=None, help="Megatron-LM source tree (optional)") + parser.add_argument("--quick", action="store_true", help="small sizes, for a smoke test only") + parser.add_argument("--gate-child", nargs=2, default=None, help=argparse.SUPPRESS) + args = parser.parse_args() + if args.megatron_src: + sys.path.insert(0, args.megatron_src) + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + if args.gate_child: + _gate_child(args.gate_child[0], args.op, args.gate_child[1]) + return + if args.op == "rms_norm_gated": + from rl_engine.backends.cuda.norm import rmsnorm as R + + if not hasattr(R, "Qwen3NextRMSNormGatedCudaOp"): + raise SystemExit("rms_norm_gated: the gated CUDA op is not on this branch") + torch.backends.cuda.matmul.allow_tf32 = False + checks = set(args.checks.split(",")) + results: dict[str, Any] = {} + for name, cand in candidates(args.op).items(): + if "unavailable" in cand: + results[name] = cand + print(f"{name}: unavailable ({cand['unavailable'][:80]})", flush=True) + continue + entry: dict[str, Any] = {"source": cand["source"], "backward": cand["backward"]} + if "bi" in checks: + entry["batch_invariance"] = batch_invariance(args.op, cand, args.quick) + if "perf" in checks: + entry["accuracy_latency"] = accuracy_latency(args.op, cand, args.quick) + if "gates" in checks and cand["backward"] == "full" and name != "reference": + entry["contract_gates"] = contract_gates(args.op, name, args.megatron_src) + results[name] = entry + summary = {k: v for k, v in entry.items() if k != "source"} + print(f"{name}: {json.dumps(summary)[:300]}", flush=True) + torch.cuda.empty_cache() + + megatron_commit = None + if args.megatron_src: + megatron_commit = ( + subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=args.megatron_src, + capture_output=True, + text=True, + ).stdout.strip() + or None + ) + report = { + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": args.op, + "rl_kernel_commit": _git("rev-parse", "HEAD"), + "tracked_tree_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": { + "gpu": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "cuda": torch.version.cuda, + "libraries": {lib: _version(lib) for lib in LIBS}, + "megatron_source_commit": megatron_commit, + }, + "quick": args.quick, + "checks": sorted(checks), + "results": results, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + plot(report, args.out) + commit, dirty = report["rl_kernel_commit"][:7], report["tracked_tree_dirty"] + print(f"wrote {args.out} (commit {commit}, dirty={dirty})") + + +if __name__ == "__main__": + main()