diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4634569cd..e1d67c01d 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 tests/models/qwen3_next/test_gdn_state_contract.py -q - name: Run Cross-Configuration Contract Tests (CPU-safe) run: | diff --git a/.github/workflows/qwen3-next-provider-gpu.yml b/.github/workflows/qwen3-next-provider-gpu.yml new file mode 100644 index 000000000..a21750d8e --- /dev/null +++ b/.github/workflows/qwen3-next-provider-gpu.yml @@ -0,0 +1,84 @@ +# SPDX-License-Identifier: Apache-2.0 +# RFC #428 Qwen3-Next provider comparisons: the C1/C3/C4 norm providers and the C6 +# GDN decode-step goldens, each checked against the real vLLM 0.30.0 provider. +# +# These check_ files import vLLM, so the default pytest collection never runs them +# (tests/integrations/common/test_framework_operator_integrations.py asserts vLLM is not imported). This +# workflow runs them in their own process on a self-hosted runner that a maintainer +# registers with the labels below and keeps on the audited runtime (CUDA 13, +# torch 2.13.0+cu130, vLLM 0.30.0). Without such a runner the job queues rather than +# reporting a false pass. +# +# Security: do not use pull_request_target. Fork PRs never reach the self-hosted +# runner; a maintainer dispatches the reviewed commit from a trusted branch. + +name: Qwen3-Next-provider-GPU + +on: + pull_request: + branches: [main, test-qwennext] + paths: + - "rl_engine/backends/**" + - "rl_engine/reference/**" + - "rl_engine/validation/models/qwen3_next*" + - "tools/validation/models/*qwen3_next*" + - "tests/models/qwen3_next/**" + - "tests/models/qwen3_next/check_gdn_recurrent_golden.py" + - ".github/workflows/qwen3-next-provider-gpu.yml" + workflow_dispatch: + +concurrency: + group: qwen3-next-provider-gpu-${{ github.ref }} + cancel-in-progress: false + +permissions: + contents: read + +jobs: + fork-pr-notice: + if: github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository + runs-on: ubuntu-latest + steps: + - name: Report required trusted execution + run: | + echo "Fork code does not run on the self-hosted Qwen3-Next provider runner." + echo "A maintainer must dispatch this workflow from a trusted upstream branch." + echo "source_repository=${{ github.event.pull_request.head.repo.full_name }}" + echo "source_sha=${{ github.event.pull_request.head.sha }}" + + provider: + if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository + runs-on: [self-hosted, linux, x64, rl-kernel-qwen3-next] + timeout-minutes: 30 + steps: + - uses: actions/checkout@v4 + with: + persist-credentials: false + - name: Require the audited runtime + run: | + python3 - <<'PY' + import torch, vllm + assert torch.__version__ == "2.13.0+cu130" + assert vllm.__version__ == "0.30.0" + assert torch.cuda.is_available() + PY + - name: Build the extension against that runtime + run: python3 setup.py build_ext --inplace + - name: Require the extension built from this checkout + # A stale editable install elsewhere on the runner would otherwise satisfy the + # import and test some other tree's _C. + run: | + python3 - <<'PY' + import os + from pathlib import Path + + from rl_engine import _C + + workspace = Path(os.environ["GITHUB_WORKSPACE"]).resolve() + loaded = Path(_C.__file__).resolve() + assert workspace in loaded.parents, f"rl_engine._C loaded from {loaded}, outside {workspace}" + print("rl_engine._C:", loaded) + PY + - name: Check pinned providers in their own process + run: | + python3 -m pytest -q tests/models/qwen3_next/check_qwen3_next_norm_providers.py tests/models/qwen3_next/check_gdn_recurrent_golden.py 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..1c5a2be53 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 \ @@ -75,3 +76,14 @@ for cell in required: ) print("[ws1-gtest] C8 gate passed") PY + +echo "[ws1-gtest] Qwen3-Next C3/C4 norms (operator scope)" +QWEN3_NEXT_NORM_MANIFEST=rl_engine/validation/models/qwen3_next_norm_manifest.json +"$PY" tools/validation/operators/check_forward_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op qwen3_next_rms_norm --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_gradient_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op qwen3_next_rms_norm --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_forward_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op rms_norm_gated --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_gradient_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op rms_norm_gated --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 diff --git a/csrc/bindings/ops.cpp b/csrc/bindings/ops.cpp index 0c7c3e7f4..a6e0d52f0 100644 --- a/csrc/bindings/ops.cpp +++ b/csrc/bindings/ops.cpp @@ -263,14 +263,36 @@ 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_gated_forward_cuda( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor y, + torch::Tensor rstd, + double eps, + double weight_offset, + int64_t activation); + +void rmsnorm_gated_backward_dx_cuda( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + torch::Tensor dx, + double weight_offset, + int64_t activation); void rmsnorm_backward_partial_dw_cuda( torch::Tensor dy, @@ -296,23 +318,39 @@ static void rmsnorm_check_input(const torch::Tensor& x, const char* name) { TORCH_CHECK(x.is_contiguous(), name, " must be contiguous"); } +static void rmsnorm_check_weight(const torch::Tensor& x, const torch::Tensor& weight) { + TORCH_CHECK(x.dim() == 2 && x.size(1) > 0, "x must be 2D with positive hidden size"); + TORCH_CHECK(x.scalar_type() == torch::kFloat32 || x.scalar_type() == torch::kFloat16 || + x.scalar_type() == torch::kBFloat16, "x must be float32, float16 or bfloat16"); + TORCH_CHECK(weight.dim() == 1 && weight.size(0) == x.size(1), "weight must be [H]"); + TORCH_CHECK(weight.device() == x.device(), "weight must be on the same device as x"); + TORCH_CHECK(weight.scalar_type() == x.scalar_type() || weight.scalar_type() == torch::kFloat32, + "weight must have x dtype or float32"); +} + +static void rmsnorm_check_backward(const torch::Tensor& dy, const torch::Tensor& x, + const torch::Tensor& rstd) { + TORCH_CHECK(dy.device() == x.device() && rstd.device() == x.device(), + "dy and rstd must be on the same device as x"); + TORCH_CHECK(dy.scalar_type() == x.scalar_type(), "dy must have the same dtype as x"); + TORCH_CHECK(rstd.scalar_type() == torch::kFloat32, "rstd must be float32"); +} + 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"); - - TORCH_CHECK(x.dim() == 2, "x must be 2D [T, H]"); - TORCH_CHECK(weight.dim() == 1, "weight must be 1D [H]"); - TORCH_CHECK(x.size(1) == weight.size(0), "x.size(1) must equal weight.size(0)"); + rmsnorm_check_weight(x, weight); auto T = x.size(0); 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,22 +359,89 @@ 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"); rmsnorm_check_input(weight, "weight"); + rmsnorm_check_weight(x, weight); rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); TORCH_CHECK(dy.sizes() == x.sizes(), "dy and x must have same shape"); - TORCH_CHECK(x.dim() == 2, "x must be 2D [T, H]"); - TORCH_CHECK(weight.dim() == 1, "weight must be 1D [H]"); TORCH_CHECK(rstd.dim() == 1, "rstd must be 1D [T]"); TORCH_CHECK(rstd.size(0) == x.size(0), "rstd.size(0) must equal x.size(0)"); 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; +} + +static void rmsnorm_gated_check( + const torch::Tensor& x, + const torch::Tensor& weight, + const torch::Tensor& gate, + int64_t activation) +{ + rmsnorm_check_input(x, "x"); + rmsnorm_check_input(weight, "weight"); + rmsnorm_check_weight(x, weight); + rmsnorm_check_input(gate, "gate"); + TORCH_CHECK(gate.device() == x.device(), "gate must be on the same device as x"); + TORCH_CHECK(gate.sizes() == x.sizes(), "gate must have the same shape as x"); + TORCH_CHECK( + gate.scalar_type() == x.scalar_type(), + "gate must have the same dtype as x"); + // 0 = silu/swish, 1 = sigmoid. Anything else fails closed rather than + // silently computing a different activation (RFC #428 section 6, item 7). + TORCH_CHECK( + activation == 0 || activation == 1, + "activation must be 0 (silu) or 1 (sigmoid), got ", activation); +} + +std::vector rmsnorm_gated_forward( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + double eps, + double weight_offset, + int64_t activation) +{ + rmsnorm_gated_check(x, weight, gate, activation); + + auto T = x.size(0); + auto y = torch::empty_like(x); + auto rstd = torch::empty({T}, x.options().dtype(torch::kFloat32)); + + rmsnorm_gated_forward_cuda(x, weight, gate, y, rstd, eps, weight_offset, activation); + + return {y, rstd}; +} + +torch::Tensor rmsnorm_gated_backward_dx( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + double weight_offset, + int64_t activation) +{ + rmsnorm_gated_check(x, weight, gate, activation); + rmsnorm_check_input(dy, "dy"); + rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); + + TORCH_CHECK(dy.sizes() == x.sizes(), "dy must have the same shape as x"); + TORCH_CHECK(rstd.dim() == 1 && rstd.size(0) == x.size(0), "rstd must be [T]"); + + auto dx = torch::empty_like(x); + + rmsnorm_gated_backward_dx_cuda( + dy, x, weight, gate, rstd, dx, weight_offset, activation); return dx; } @@ -350,6 +455,7 @@ torch::Tensor rmsnorm_backward_dw( rmsnorm_check_input(dy, "dy"); rmsnorm_check_input(x, "x"); rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); rmsnorm_check_input(mask, "mask"); TORCH_CHECK(dy.sizes() == x.sizes(), "dy and x must have same shape"); @@ -357,6 +463,8 @@ torch::Tensor rmsnorm_backward_dw( TORCH_CHECK(rstd.dim() == 1, "rstd must be 1D [T]"); TORCH_CHECK(mask.dim() == 1, "mask must be 1D [T]"); TORCH_CHECK(mask.scalar_type() == torch::kBool, "mask must be bool"); + TORCH_CHECK(mask.device() == x.device(), "mask must be on the same device as x"); + TORCH_CHECK(x.size(1) > 0, "hidden size must be positive"); TORCH_CHECK(rstd.size(0) == x.size(0), "rstd.size(0) must equal x.size(0)"); TORCH_CHECK(mask.size(0) == x.size(0), "mask.size(0) must equal x.size(0)"); @@ -718,9 +826,24 @@ 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"); + m.def("rmsnorm_gated_forward", &rmsnorm_gated_forward, + "Batch-invariant gated RMSNorm forward CUDA", + py::arg("x"), py::arg("weight"), py::arg("gate"), py::arg("eps"), + py::arg("weight_offset") = 0.0, py::arg("activation") = 0); + m.def("rmsnorm_gated_backward_dx", &rmsnorm_gated_backward_dx, + "Batch-invariant gated RMSNorm backward dx CUDA", + py::arg("dy"), py::arg("x"), py::arg("weight"), py::arg("gate"), + py::arg("rstd"), py::arg("weight_offset") = 0.0, + py::arg("activation") = 0); #if !defined(USE_ROCM) m.def( "reduce_rows_fp32_left_fold", diff --git a/csrc/cuda/norm/rmsnorm.cu b/csrc/cuda/norm/rmsnorm.cu index b32bc5af4..81a02cd2d 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); @@ -187,6 +198,133 @@ __global__ void rmsnorm_bwd_dx_kernel( } +// --------------------------------------------------------------------------- +// Gated RMSNorm (Qwen3-Next GDN block). +// +// y = x * rstd * (weight_offset + weight) * act(gate) +// +// The gate activation and the weight multiply are both evaluated in fp32 and +// there is exactly one cast, at the store. This mirrors vLLM's RMSNormGated +// with norm_before_gate=True, which is how the GDN block constructs it. +// +// ACT selects the gate activation: 0 = silu/swish, 1 = sigmoid. expf (not the +// __expf intrinsic) is used deliberately -- the fast intrinsic trades accuracy +// for speed and would put the result further from the fp32 reference. +// --------------------------------------------------------------------------- + +template +__device__ __forceinline__ float gate_activation(float z) { + const float sigma = 1.0f / (1.0f + expf(-z)); + return (ACT == 0) ? z * sigma : sigma; +} + +template +__device__ __forceinline__ float gate_activation_grad(float z) { + const float sigma = 1.0f / (1.0f + expf(-z)); + // d/dz [z * sigma] = sigma * (1 + z * (1 - sigma)); d/dz [sigma] = sigma * (1 - sigma) + return (ACT == 0) ? sigma * (1.0f + z * (1.0f - sigma)) : sigma * (1.0f - sigma); +} + + +template +__global__ void rmsnorm_gated_fwd_kernel( + const scalar_t* __restrict__ x, + const weight_t* __restrict__ weight, + const scalar_t* __restrict__ gate, + scalar_t* __restrict__ y, + float* __restrict__ rstd, + int T, + int H, + float eps, + float weight_offset +) { + int row = blockIdx.x; + int tid = threadIdx.x; + + const scalar_t* x_row = x + row * H; + const scalar_t* gate_row = gate + row * H; + scalar_t* y_row = y + row * H; + + float local_sum = 0.0f; + + // The statistic is over x only; the gate never enters the reduction, so + // rstd here is bit-identical to the ungated kernel's for the same x. + for (int col = tid; col < H; col += blockDim.x) { + float xv = load_as_float(x_row + col); + local_sum += xv * xv; + } + + float sum = block_reduce_sum(local_sum); + + float row_rstd = rsqrtf(sum / static_cast(H) + eps); + + if (tid == 0) { + rstd[row] = row_rstd; + } + + __syncthreads(); + + for (int col = tid; col < H; col += blockDim.x) { + float xv = load_as_float(x_row + col); + float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; + float zv = load_as_float(gate_row + col); + float out = xv * row_rstd * wv * gate_activation(zv); + store_from_float(y_row + col, out); + } +} + + +template +__global__ void rmsnorm_gated_bwd_dx_kernel( + const scalar_t* __restrict__ dy, + const scalar_t* __restrict__ x, + const weight_t* __restrict__ weight, + const scalar_t* __restrict__ gate, + const float* __restrict__ rstd, + scalar_t* __restrict__ dx, + int T, + int H, + float weight_offset +) { + int row = blockIdx.x; + int tid = threadIdx.x; + + const scalar_t* dy_row = dy + row * H; + const scalar_t* x_row = x + row * H; + const scalar_t* gate_row = gate + row * H; + scalar_t* dx_row = dx + row * H; + + float local_dot = 0.0f; + + // Identical to the ungated dx, with the per-column scale (w + offset) + // replaced by (w + offset) * act(gate): the gate is a constant wrt x. + for (int col = tid; col < H; col += blockDim.x) { + 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 zv = load_as_float(gate_row + col); + local_dot += dyv * (wv * gate_activation(zv)) * xv; + } + + float dot = block_reduce_sum(local_dot); + + float r = rstd[row]; + float coeff = dot * r * r * r / static_cast(H); + + for (int col = tid; col < H; col += blockDim.x) { + 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 zv = load_as_float(gate_row + col); + + float out = r * dyv * (wv * gate_activation(zv)) - xv * coeff; + store_from_float(dx_row + col, out); + } +} + template __global__ void rmsnorm_partial_dw_kernel( const scalar_t* __restrict__ dy, @@ -247,12 +385,14 @@ 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. const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); + if (T == 0) return; int H = x.size(1); int threads = choose_threads(H); size_t smem = threads * sizeof(float); @@ -270,10 +410,12 @@ void rmsnorm_forward_cuda( rstd.data_ptr(), T, H, - static_cast(eps) + static_cast(eps), + static_cast(weight_offset) ); }); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } @@ -282,10 +424,12 @@ 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); + if (T == 0) return; int H = x.size(1); int threads = choose_threads(H); size_t smem = threads * sizeof(float); @@ -303,13 +447,94 @@ void rmsnorm_backward_dx_cuda( rstd.data_ptr(), dx.data_ptr(), T, - H + H, + static_cast(weight_offset) ); }); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + + +void rmsnorm_gated_forward_cuda( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor y, + torch::Tensor rstd, + double eps, + double weight_offset, + int64_t activation +) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); + int T = x.size(0); + if (T == 0) return; + int H = x.size(1); + int threads = choose_threads(H); + size_t smem = threads * sizeof(float); + + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, x.scalar_type(), "rmsnorm_gated_forward_cuda", [&] { + using x_t = scalar_t; + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, weight.scalar_type(), "rmsnorm_gated_forward_weight_cuda", [&] { + using w_t = scalar_t; + if (activation == 0) { + rmsnorm_gated_fwd_kernel<<>>( + x.data_ptr(), weight.data_ptr(), gate.data_ptr(), + y.data_ptr(), rstd.data_ptr(), + T, H, static_cast(eps), static_cast(weight_offset)); + } else { + rmsnorm_gated_fwd_kernel<<>>( + x.data_ptr(), weight.data_ptr(), gate.data_ptr(), + y.data_ptr(), rstd.data_ptr(), + T, H, static_cast(eps), static_cast(weight_offset)); + } + }); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } +void rmsnorm_gated_backward_dx_cuda( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + torch::Tensor dx, + double weight_offset, + int64_t activation +) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); + int T = x.size(0); + if (T == 0) return; + int H = x.size(1); + int threads = choose_threads(H); + size_t smem = threads * sizeof(float); + + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, x.scalar_type(), "rmsnorm_gated_backward_dx_cuda", [&] { + using x_t = scalar_t; + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, weight.scalar_type(), "rmsnorm_gated_backward_dx_weight_cuda", [&] { + using w_t = scalar_t; + if (activation == 0) { + rmsnorm_gated_bwd_dx_kernel<<>>( + dy.data_ptr(), x.data_ptr(), weight.data_ptr(), + gate.data_ptr(), rstd.data_ptr(), dx.data_ptr(), + T, H, static_cast(weight_offset)); + } else { + rmsnorm_gated_bwd_dx_kernel<<>>( + dy.data_ptr(), x.data_ptr(), weight.data_ptr(), + gate.data_ptr(), rstd.data_ptr(), dx.data_ptr(), + T, H, static_cast(weight_offset)); + } + }); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + void rmsnorm_backward_partial_dw_cuda( torch::Tensor dy, torch::Tensor x, @@ -319,6 +544,7 @@ void rmsnorm_backward_partial_dw_cuda( ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); + if (T == 0) return; int H = x.size(1); int chunks = (T + RMSNORM_DW_ROWS_PER_CHUNK - 1) / RMSNORM_DW_ROWS_PER_CHUNK; @@ -338,6 +564,7 @@ void rmsnorm_backward_partial_dw_cuda( H ); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } @@ -360,6 +587,7 @@ void rmsnorm_backward_reduce_dw_cuda( chunks, H ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } #if !defined(USE_ROCM) diff --git a/docs/.nav.yml b/docs/.nav.yml index 29b5ead37..4e04fd061 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -29,6 +29,8 @@ nav: - operators/sampling.md - operators/det-gemm.md - operators/embedding.md + - operators/qwen3-next-rms-norm.md + - operators/qwen3-next-rms-norm-gated.md - Developer Guide: - contributing/README.md - Contributor Guide: contributing/contributor-guide.md diff --git a/docs/archive/design/rfc428-c6-gdn-recurrent-replay.md b/docs/archive/design/rfc428-c6-gdn-recurrent-replay.md new file mode 100644 index 000000000..84a977547 --- /dev/null +++ b/docs/archive/design/rfc428-c6-gdn-recurrent-replay.md @@ -0,0 +1,303 @@ +# RFC #428 C6 — Qwen3-Next Gated DeltaNet recurrent replay + +Design notes for RFC #428 work item C6 (GDN recurrent response replay) on the CUDA +track. Measured on 2× B200 (sm_100), torch 2.13.0+cu130, vllm 0.30.0, +transformers 5.17.0. vLLM paths below are relative to the installed `vllm` 0.30.0 +package; two files are named `causal_conv1d.py`, and each citation says whether it +means vLLM's (`model_executor/layers/mamba/ops/causal_conv1d.py`) or the golden's +(`rl_engine/reference/linear_attn/causal_conv1d.py`). + +Claim level reached: **L0 repeatable, L1 batch-invariant**. L2 is not claimed. + +## 1. Which provider a rollout decode actually takes + +vLLM 0.30.0 picks the GDN decode path per engine step. The choice depends on env +defaults, on what else the step contains, and on how the model constructs the layer: + +| step contains | conv | recurrence | +|---|---|---| +| decodes only, no draft tokens | `causal_conv1d_update` | `fused_recurrent_gated_delta_rule_packed_decode` | +| decodes and at least one prefill, no draft tokens | `causal_conv1d_fn` | `fused_sigmoid_gating_delta_rule_update` for the cached decode rows | +| draft tokens (speculative decode / MTP) | `causal_conv1d_update` with `num_accepted_tokens` | `fused_sigmoid_gating_delta_rule_update`; plain decodes in the step are reclassified as prefills | + +- **Decode-only.** With `VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE` at its default of on + (`build_tools/envs.py:1199-1200`), a decode-only step returns early at + `model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py:1295-1307` into + `_forward_core_decode_non_spec`, which calls the packed kernel (`:1697-1708`). **The + goldens target this path.** With the variable off, the same step goes through + `fused_sigmoid_gating_delta_rule_update` instead (`:1553-1572`). +- **Mixed decode and prefill.** `_forward_core` runs the conv for the whole non-spec + batch through `causal_conv1d_fn` (branch at `:1373`, call at `:1378-1388`) and the + cached decode rows through `fused_sigmoid_gating_delta_rule_update` + (`split_non_spec` defined at `:1408-1412`, branch at `:1493`, call at `:1497-1512`). + See §5 for what that does to the claim. +- **Draft tokens.** Spec rows go through `causal_conv1d_update` with + `num_accepted_tokens` (`:1357-1370`) and `fused_sigmoid_gating_delta_rule_update` + (`:1470-1487`). Plain decodes in the same step are reclassified as prefills + (`v1/attention/backends/gdn_attn.py:283-289`). A step with speculative decoding + enabled but zero draft tokens sets `spec_sequence_masks` to `None` + (`gdn_attn.py:236-243`) and is treated as decode-only. + +`VLLM_GDN_DECODE_KERNEL` defaults to `"cuda"`, but that is not what Qwen3-Next runs. +`model_executor/models/qwen3_next.py:495` constructs the layer with +`gqa_interleaved_layout=True`, which makes `_fused_gdn_decode_unsupported_reason` +(`qwen_gdn_linear_attn.py:535-554`) return a reason. The layer then logs a fallback +to `"triton"` (`:520-523`), or raises `ValueError` if `VLLM_GDN_DECODE_KERNEL` was +explicitly set to `cuda` (`:516-519`). As a consequence the fused +`torch.ops._C.fused_gdn_decode_post_conv_mtp` path is unreachable for Qwen3-Next: +`_can_use_fused_gdn_mtp_decode` requires `gdn_decode_kernel == "cuda"` (`:1834`). +This section is read from the source. + +`tests/models/qwen3_next/check_qwen3_next_norm_providers.py` asserts both env defaults (`:129-142`). For +`VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE` that is a useful guard: a vLLM bump that +flips it fails loudly. For `VLLM_GDN_DECODE_KERNEL` it guards nothing for Qwen3-Next, +because the interleaved layout, not the default, decides the kernel. Neither +assertion checks which kernel actually runs. + +## 2. How the golden relates to the kernel + +The golden follows the kernel's arithmetic, not `modeling_qwen3_next.py`'s. Checked +line by line against `third_party/flash_linear_attention/ops/fused_recurrent.py:288-335`, +three places where the kernel and the HF model differ: + +1. **The gating is fused, in fp32.** `beta = sigmoid(b)` and + `g = -exp(A_log) * softplus(a + dt_bias)` are computed inside the Triton kernel, + with a `softplus` threshold branch at 20. HF computes them as separate PyTorch + ops — a different rounding path. +2. **No `repeat_interleave`.** The kernel indexes `i_h = i_hv // (HV // H)`, so q/k + stay at 16 heads while v has 32. HF materializes the repeat. +3. **The QK norm is an L2 norm over a plain sum**, `x / sqrt(sum(x*x) + 1e-6)` — not + an RMSNorm, not `F.normalize`, and dividing by `sqrt` rather than multiplying by + `rsqrt`, which differs in the last bit. `scale` is applied to `q` *after* the + norm; `k` is never scaled. + +Where the golden is **not** a transcription: + +- softplus: the kernel computes `tl.log(1.0 + tl.exp(x))` (`fused_recurrent.py:327`); + the golden computes `torch.log1p(torch.exp(safe))`, where `safe` is `x` on the taken + branch (`gated_delta_rule.py:110-112`). These round differently. +- The kernel's `exp`/`log` become `fast_expf`/`fast_logf` when `FLA_USE_FAST_OPS=1` + (`third_party/flash_linear_attention/ops/op.py:16-25`). The golden models the + default. + +Checked and found **not** to be a divergence: prefill passes +`use_qk_l2norm_in_kernel=False` (`qwen_gdn_linear_attn.py:1542`) only because +`fused_post_conv_prep(apply_l2norm=True)` (`:1450`) already normalized q/k. Decode +passes `True` (`:1707`). Both paths normalize exactly once. + +## 3. State ABI + +Mirrored rather than reinvented: + +- recurrent state `[num_blocks, HV, V, K]`, V-major, addressed by `ssm_state_indices` +- conv state `[num_blocks, dim, width-1]`, layout chosen by the global + `is_conv_state_dim_first()`; the golden takes it as an argument and both are tested +- the accumulator is fp32 for the whole step; the store rounds to the state tensor's + dtype, which `FUSED_GDN_STATE_DTYPES` (`qwen_gdn_linear_attn.py:91`) allows to be + fp32 **or** bf16 +- `causal_conv1d_update` casts `x` to the cache dtype before computing + +`NULL_BLOCK_ID` is **not** one contract across the two providers: + +| | recurrent provider | conv provider | both goldens | +|---|---|---|---| +| which indices skip | `<= 0` (`fused_recurrent.py:300`) | `== null_block_id` only, default `NULL_BLOCK_ID` = 0 (vLLM `causal_conv1d.py:835`; `v1/attention/backends/utils.py:47`) | `<= 0` | +| output for a skipped row | zeros (`fused_recurrent.py:301-302`) | not written (vLLM `causal_conv1d.py:835-839` returns before any store) | zeros | +| state block | not touched | not touched | not touched | + +A negative conv index is therefore inactive in the golden +(`tests/models/qwen3_next/test_gdn_state_contract.py:54-59`) but a real, out-of-range index to the conv +provider. + +Contractions use `_chunked_sum` (`gated_delta_rule.py:76-90`): fixed 32-wide chunks, +so the reduction shape per row does not depend on the batch size, rather than +`torch.matmul`, whose reduction order is unspecified. That is the argument for the +L1 claim; L1 itself is established empirically by `test_golden_is_batch_invariant`. +The same choice is why the golden is not bitwise against the kernel's reduction. + +## 4. Agreement with the provider + +Qwen3-Next dims (H=16, HV=32, K=V=128), bf16 I/O, `use_qk_l2norm_in_kernel=True`, +random inputs, B ∈ {1, 4, 17, 64}, one seed per batch. Bounds asserted by +`tests/models/qwen3_next/check_gdn_recurrent_golden.py` (`_RECURRENT_BOUNDS`). They, and every other +bound in that file, are **regression bounds against the provider, not gate evidence**; +they do not go through `resolve_tolerance`. For scale, the gate contract's +`forward_accuracy/by_op_class/reduction` row in +`rl_engine/contracts/profiles/precision/ws1.json` is atol = rtol = 1e-4 for float32 +and atol = 5e-2, rtol = 2e-2 for bfloat16. + +| | max\|diff\| out | max\|diff\| state | +|---|---|---| +| fp32 state | ≤ 1e-3 | ≤ 1e-5 | +| bf16 state | ≤ 1e-3 | ≤ 5e-3 | + +Measured values for the same inputs, from `tools/validation/models/ws1_gdn_provider_agreement.py` on +B200 at commit `acf38b6` (clean checkout). The last column counts output elements +whose bits differ: + +| state | B | max\|diff\| out | max\|diff\| state | out elements differing | +|---|---|---|---|---| +| fp32 | 1 | 1.49e-08 | 1.19e-07 | 1 / 4096 | +| fp32 | 4 | 9.54e-07 | 1.79e-07 | 2 / 16384 | +| fp32 | 17 | 3.81e-06 | 2.38e-07 | 10 / 69632 | +| fp32 | 64 | 6.10e-05 | 2.98e-07 | 40 / 262144 | +| bf16 | 1 | 3.73e-09 | 9.77e-04 | 2 / 4096 | +| bf16 | 4 | 1.53e-05 | 9.77e-04 | 2 / 16384 | +| bf16 | 17 | 3.05e-05 | 1.95e-03 | 14 / 69632 | +| bf16 | 64 | 3.05e-05 | 1.95e-03 | 36 / 262144 | + +This reproduces the figures an earlier version of this note quoted without a runner +(out 1.5e-08 .. 6.1e-05 and state ≤ 3.0e-07 for fp32; out 3.7e-09 .. 3.1e-05 and +state ≤ 2.0e-03 for bf16). It is one seed per batch on one device. + +Causal conv uses sequential FP32 accumulation **starting from bias**, with +products first rounded to the operand dtype (golden `causal_conv1d.py:167-177`; +the provider initialises from bias at vLLM `causal_conv1d.py:960-967`, `1000`). The +previous BF16 path incorrectly promoted both operands to FP32; its disagreements +were not limited to one BF16 ULP. The CPU tests in `tests/models/qwen3_next/test_gdn_state_contract.py` +cover bias order and bf16 product rounding, and +`test_conv_provider_preserves_bf16_product_cancellation` checks the provider on a +constructed cancellation input. Provider comparisons limit mismatches to 32 elements +on the checked fixtures (a regression bound, not gate evidence). This is not a bitwise claim. + +Measured with the same runner and commit: the rolled conv state is bitwise equal in +all 8 (batch, cache dtype) cases. Output elements differing, with an fp32 cache: +0 / 8192 (B=1), 0 / 32768 (B=4), 1 / 139264 (B=17, max|diff| 2.44e-04) and +5 / 524288 (B=64, max|diff| 3.91e-03). With a bf16 cache: 0 at B=1, 4 and 17, and +3 / 524288 at B=64 (max|diff| 1.56e-02). + +**Why a few fp32-cache outputs differ.** There are two mechanisms, both on the provider +side, and which one applies depends on the Triton specialization. Measured by +`tools/validation/models/ws1_gdn_provider_agreement.py` (`conv_silu`, `conv_noact_bf16`) on B200 at commit +`bb89750`: + +*With SiLU -- Qwen3-Next's configuration and the check file's -- the activation's +implementation.* + +- With the activation off and the output kept in fp32, provider and golden agree bitwise + on every element at B = 1, 4, 17 and 64: the cast to the cache dtype, the bias start, + the product rounding and the tap order all match. +- Both sides compute `acc / (1 + exp(-acc))` (vLLM `causal_conv1d.py:1085`), but Triton + lowers the exp and the fp32 division to `ex2.approx` and `div.full.f32`, while the + golden's PyTorch result equals the same expression with `libdevice.exp` and IEEE + division (`div_rn`). Applying Triton's `x / (1 + tl.exp(-x))` to the golden's + pre-activation values reproduces the provider's fp32 output bitwise; replacing only the + exp, or only the division, reproduces neither side. +- In fp32 about 38% of outputs differ: 86% of those by 1 ULP, 12% by 2, 3% by 3 or 4, + 0.2% by more. Only values on opposite sides of a bf16 rounding midpoint survive the bf16 + store: 0, 0, 1 and 5 elements, each one bf16 ULP apart. +- FP contraction plays no part on this path. These specializations compute the tap + products as packed `mul.f32x2` with separate adds and contain no `FFMA`; recompiling + with `TRITON_DEFAULT_FP_FUSION=0` leaves every count unchanged. + +*Without an activation and with a bf16 output -- FP contraction.* + +- With fusion at its default, this specialization contracts each tap's multiply-add + (`fma.rn.f32x2` in the PTX, `FFMA` in the SASS). Its fp32-output twin computes the same + values (the provider casts x to the fp32 cache dtype first) but is not contracted, and + matches the golden bitwise. +- The bf16 outputs differ in 1, 0, 4 and 14 elements (max |diff| 3.9e-3). 18 of the 19 + have |out| < 0.11, where the taps nearly cancel; there the gap can span several bf16 + ULPs (up to 96 at |out| ≈ 2.4e-7) while staying at most 2.4e-7 in absolute terms. +- With `TRITON_DEFAULT_FP_FUSION=0` this specialization compiles to separate multiplies + and adds, and the mismatches drop to 0. The golden is self-consistent across output + dtypes. Qwen3-Next's conv applies SiLU, so this specialization is not on its decode + path. + +*Open:* the recurrent provider computes `exp(g)` and `sigmoid(b)` with Triton (`tl.exp` +via the vendored FLA `op.py`), the golden with PyTorch; the same kind of difference may +account for part of the recurrent output mismatches in the table above. Not +investigated. + +## 5. Decode versus chunked prefill + +The earlier 1024-step drift and prefill tables did not have a checked-in runner; +they are withdrawn as acceptance evidence. The existing 128-step synthetic test +only bounds its fixed seed and gate inputs. It does not establish a universal +plateau, prefill/decode equality, or any bound on model logits. + +**Mixed decode-and-prefill steps (provider side).** The goldens target the +decode-only path (§1). In a step that also holds a prefill, vLLM 0.30.0 sends the +cached decode rows through `causal_conv1d_fn` and +`fused_sigmoid_gating_delta_rule_update` instead, so in those steps the recurrent +golden does not target what rollout runs. + +*The numbers in the rest of this subsection are a reported observation with no +checked-in runner in this repository. By the same rule that withdrew the tables above, +they are not acceptance evidence.* + +A measurement on B200 drove the real `GDNAttentionMetadataBuilder.build()` and +`_forward_core` of **one standalone layer** built from the checkpoint's +`config.json`. Its limits: parameters and cache were synthetic (no weights loaded); it +ran in a single process, with no engine, scheduler or CUDA graph; and the test, not the +scheduler, built the attention metadata (decode rows first). For one target decode +request with an fp32 recurrent state, three prefill-bearing step compositions gave +identical numbers: + +| head counts | step's bf16 output | fp32 state | +|---|---|---| +| TP1 (H=16, HV=32) | matched | 75,146 of 524,288 elements differ (8.8e-08 relative) | +| TP4 per-rank (H=4, HV=8) | 1 element differs (4.8e-05 relative) | 32,663 of 131,072 differ (1.7e-07) | + +The convolution output and conv state matched bitwise in every composition, so the +difference comes from the recurrent kernels, not from `causal_conv1d_fn`. At TP1 the +state difference reached a bf16 output within the next 16 plain decode steps in 4 of +20 seeds; that was measured at TP1 only, and with synthetic parameters the frequency +does not carry over to real weights. At TP4 per-rank head counts it appears in the +same step's output. This is a provider-side batch-composition dependence: whether a +prefill shares the step changes a decode row's state, and at TP4 per-rank head counts +its output too, which no golden can fix. The measurement was made outside this change +and is not yet published. + +vllm-project/vllm#49827, open and unmerged, routes mixed-step *recurrent* decodes +through the packed kernel too. Its two commits, applied to 0.30.0, closed the gap to 0 +at both head counts above; its scheduler part was not tested. It does not change the +conv path in mixed steps, which stays `causal_conv1d_fn`, and its own validation is on +Qwen3.5 (non-interleaved), H100, TP1. Disabling +`VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE` instead sends every non-speculative decode +row through `fused_sigmoid_gating_delta_rule_update`; the same measurement found the +decode row's output and state bitwise equal in every step composition in that +configuration, and the recurrent golden's target would then have to change. + +The strict profile uses FP32 recurrent state. BF16 state remains a differential +experiment. Full checkpoint prefill, response replay, optimizer updates and +reload all remain required before L2 can pass. + +Cache indices must be int32/int64, on the input device, and positive active +indices must be unique and in range. Nonpositive sentinels may repeat. The golden +validates before any cache update. Raw provider calls remain the caller's +responsibility: neither provider checks index values. `causal_conv1d_update`'s +`validate_data=True` adds only shape and stride asserts and a +`null_block_id is not None` assert (vLLM `causal_conv1d.py:1151-1153`, `1188-1201`), +so an out-of-range index is an unchecked memory access in either mode (read from +the source, not exercised). + +## 6. Current boundary + +Not covered, with reasons: + +| deferred | why | +|---|---| +| Speculative decode / MTP | RFC #428 §2.2 excludes speculative decoding from the first claim. Qwen3-Next's MTP head ships in the checkpoint (`mtp.*`, loaded by `model_executor/models/qwen3_next_mtp.py`) and is full attention (`qwen3_next_mtp.py:90-92`). Enabling it changes the target model's GDN path in steps that carry draft tokens (§1); `fused_gdn_decode_post_conv_mtp` is unreachable for Qwen3-Next | +| Backward for the recurrent step | RFC #428 §2.2 item 4: backward need not match rollout, only be correct for the replayed forward. RFC §9.1 makes it a separate work item, **RFC #428 C7** (GDN backward/recompute adapter including prompt-state gradient). `supports_backward=false` here; `_softplus`'s NaN-gradient fix (`08969ac`) adds no backward claim | +| Provider bridge / registry entry | the golden should survive a drift sweep against a real checkpoint first | +| Paged block allocation policy | the ABI is mirrored; the allocator is not modelled | +| TP sharding of `A_log` / `dt_bias` | single card only | +| ROCm / Ascend | CUDA first, per the WS1 order | + +`runtime_verified=false` (no checkpoint), `supports_backward=false`, +`checkpoint=absent`. Op-level agreement says nothing about 48 composed layers. + +## 7. Accepted execution boundaries + +Use a shared, explicitly pinned vLLM-compatible forward provider on both sides, +with independent VIME recomputation. Disable MTP and prefix reuse. Keep FP32 +recurrent state for strict acceptance; BF16 is experimental. Require bitwise +logits/logprobs at a common topology, at least two real optimizer updates, weight +synchronization and checkpoint reload. No operator-level test closes these gates. + +The 2026-09-30 real-checkpoint startup attempt with vLLM 0.30.0 failed before +inference: `VLLM batch_invariant mode is not supported for GDN_ATTN`. A shared +provider integration must resolve this; disabling the check is not L2 evidence. *This +failure is a reported observation: the attempt's log and launcher are not checked into +this repository, so it is not acceptance evidence either.* diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..0b93b3d90 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -32,4 +32,6 @@ 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) +- [Qwen3-Next Gated RMSNorm](qwen3-next-rms-norm-gated.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/qwen3-next-rms-norm-gated.md b/docs/operators/qwen3-next-rms-norm-gated.md new file mode 100644 index 000000000..217fa45ed --- /dev/null +++ b/docs/operators/qwen3-next-rms-norm-gated.md @@ -0,0 +1,164 @@ +# Qwen3-Next Gated RMSNorm + +## Summary + +The norm inside Qwen3-Next's Gated DeltaNet block: + +``` +out = x * rstd * weight * silu(gate) +``` + +It is a different operator from the decoder norm, not a variant of it: the weight is +plain rather than zero-centred (upstream initializes it to ones), it takes a second +input tensor, and — critically — `transformers` and vLLM disagree about where the +weight multiply happens. See [Qwen3-Next RMSNorm](qwen3-next-rms-norm.md) for the +decoder-norm convention. + +Added for RFC #428 C1 on the Qwen3-Next rollout-vs-replay path. + +## Entry Point + +```python +from rl_engine.reference.norm.qwen3_next_rms_norm import ( + Qwen3NextRMSNormGatedOp, # strict / vLLM convention + Qwen3NextRMSNormGatedHFOp, # transformers convention, kept as a witness +) +from rl_engine.backends.cuda.norm.rmsnorm import ( + Qwen3NextRMSNormGatedCudaOp, + rmsnorm_gated_cuda, +) + +y = Qwen3NextRMSNormGatedCudaOp().forward(x, weight, gate, eps=1e-6) +y = Qwen3NextRMSNormGatedCudaOp(activation="sigmoid").forward(x, weight, gate, eps=1e-6) +y = rmsnorm_gated_cuda(x, weight, gate, eps=1e-6, activation="sigmoid") +``` + +## Backends + +| Backend | Wrapper | Native symbol | Status | +| --- | --- | --- | --- | +| CUDA | `Qwen3NextRMSNormGatedCudaOp` | `rl_engine._C.rmsnorm_gated_forward` | Supported | +| ROCm | — | — | Not implemented | +| PyTorch fallback | `Qwen3NextRMSNormGatedOp` | — | Supported (WS1 gold) | + +The CUDA op is deliberately **not** a subclass of `RMSNormCudaOp`: it takes an extra +required tensor, so it cannot stand in for one. + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `x` | `[..., H]` | fp32 / bf16 / fp16 | H > 0; wrapper makes contiguous copies | +| `weight` | `[H]` | matches `x` or fp32 | plain, NOT zero-centred | +| `gate` | same as `x` | same as `x` | shape, dtype and device must match before flattening | +| `eps` | scalar | float | `1e-6` for Qwen3-Next | +| `activation` | — | `"silu"`/`"swish"`/`"sigmoid"` | constructor argument; `"swish"` is an alias for `"silu"`, as in vLLM; anything else is rejected | + +All tensors must share a device. Low-level CUDA bindings require contiguous +2-D inputs, and validate backward gradient dtype and FP32 statistics. Empty +batches are supported by the bindings. + +Only `norm_before_gate=True` and `group_size=None` are implemented — the +configuration vLLM's GDN block constructs. The op has no parameter for either, so +other configurations are not implemented rather than rejected at runtime. + +## Dispatch Behavior + +Registered as `rms_norm_gated`. CUDA prefers the kernel; every other platform +resolves to the PyTorch reference. `__init__` validates the compiled symbols, so on a +build without the extension the registry falls through instead of returning an op +that raises at call time. + +## Accuracy + +Claim levels: + +- **L1** for `Qwen3NextRMSNormGatedCudaOp`: prefix slices in + `tests/models/qwen3_next/test_qwen3_next_norm.py`, and batch size, chunking, padding and permutation + in the C3/C4 gates (`ci/scripts/run_ws1_gtest.sh`). +- **L0** is not separately tested for the CUDA op; no test repeats it. +- L2 is not claimed. + +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. + +`transformers` casts the normalized value back to the input dtype *before* the weight +multiply; vLLM keeps it in fp32. `Qwen3NextRMSNormGatedOp` follows vLLM and +`Qwen3NextRMSNormGatedHFOp` keeps the transformers convention as a witness. +`test_gated_conventions_diverge_in_low_precision` asserts that the two differ in bf16 +(by more than `1e-3`) and agree bitwise in fp32. One earlier measurement, with no +committed script, found them differing on 35% of elements with `max|diff| = 6.25e-2` +(one seed, B200, bf16, `head_v_dim=128`, 512 rows, direct eager calls). That is a +one-off observation, not an assertion. Which convention the strict profile should use +is an open question for RFC #428. + +"Bitwise equal to vLLM" is undefined until a provider is named. Over 40 seeds (bf16, +`head_v_dim=128`, 512 rows, B200), vLLM's eager `forward_native` and `forward_cuda` +disagreed on 21, worst `1.56e-2`; in fp32 they stay within the asserted `1e-5`. These +figures come from `tests/models/qwen3_next/check_qwen3_next_norm_providers.py`, which imports real vLLM +and must be run explicitly. The PyTorch reference reproduces the convention, not +vLLM's reduction tree. + +`rstd` is bitwise identical to the ungated kernel's for the same `x`, asserted over +fp32/fp16/bf16 × H ∈ {128, 2048, 5120} × {silu, sigmoid} × offset ∈ {0, 1}, so a gate +leaking into the statistic would fail. + +Backward is assembled from deterministic pieces: `dx` from a row-local kernel, +`dweight` from fp32 row contributions reduced by the ascending-row left fold, and +`dgate` elementwise in fp32 with no reduction. + +## Performance Notes + +Reuses the existing `block_reduce_sum` / `choose_threads(H)` reduction, so the gate +costs one extra load and one fp32 multiply per element. + +```bash +python tools/validation/operators/check_operator.py --op rms_norm_gated --candidate cuda \ + --device cuda --dtype bf16 --check-grad +``` + +## Evidence + +![gated RMSNorm vs existing implementations on B200](../usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png) + +[`report.json`](../usage/evidence/qwen3-next-rms-norm-gated-b200/report.json) was written by +`tools/validation/models/qwen3_next_norm_evidence.py` from a clean tree at `822b085`, on an otherwise idle +B200 (torch 2.13.0+cu130, transformers 5.17.0, vLLM 0.30.0). Head dim 128, BF16; the FP64 +golden uses vLLM's convention. The same report re-measures the zero-centred op +([figure](../usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png)). + +| | rl-kernel CUDA | PyTorch reference | transformers (cast-first) | vLLM `RMSNormGated`* | +|---|---|---|---|---| +| rows differing alone vs in a batch, forward / `dx`,`dgate` (of 768) | **0 / 0** | 0 / 0 | 0 / 0 | 0 / n/a | +| forward elements correctly rounded vs the golden | **99.999%** | 99.999% | 65.6% | 99.999% | +| forward, 262144 rows | 235 µs | 774 µs | 697 µs | **81 µs** | +| backward, 262144 rows | 114 ms | **1.4 ms** | 1.5 ms | n/a | + +\* forward only, no backward. + +- **Every implementation is row-invariant** for this op. +- **transformers computes a different function:** it casts to BF16 before the weight + multiply, so 34% of its forward elements differ from vLLM's convention, and its gradients + are about twice as far from the golden. +- **The backward is about 75× slower than transformers at 262144 rows**, for the same reason + as the zero-centred op: `dweight` (128 columns here) is folded over the rows in ascending + order, one thread per column, so that it meets the gradient-invariance contract's + singleton-aggregate check bitwise. + +## Tests + +```bash +python -m pytest tests/models/qwen3_next/test_qwen3_next_norm.py -v +# imports real vLLM, so it is not collected by default: +python -m pytest tests/models/qwen3_next/check_qwen3_next_norm_providers.py -v +``` + +## Known Limitations + +- CUDA only; no ROCm, Ascend or Triton backend. +- `norm_before_gate=False` and grouped norms are not implemented. +- Not bitwise against either vLLM path (see Accuracy). +- Measured on sm_100 (B200); RFC #428 §2.2 forbids carrying the claim across + H100/H200/B100/B200. diff --git a/docs/operators/qwen3-next-rms-norm.md b/docs/operators/qwen3-next-rms-norm.md new file mode 100644 index 000000000..2fbd8baa5 --- /dev/null +++ b/docs/operators/qwen3-next-rms-norm.md @@ -0,0 +1,145 @@ +# 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; see +[Gated RMSNorm](qwen3-next-rms-norm-gated.md). + +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 path requires contiguous | +| `weight` | `[H]` | matches `x` | zero-centred (upstream inits to zeros) | +| `eps` | scalar | float | inside the sqrt; `1e-6` for Qwen3-Next | + +## Dispatch Behavior + +Registered as the `qwen3_next_rms_norm` gtest operator. On CUDA the registry +prefers `Qwen3NextRMSNormCudaOp`; every other platform resolves to the PyTorch +reference. The CUDA op validates the compiled symbols in `__init__`, so on a build +without the extension construction raises and the registry falls through to the +reference rather than 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. + +## Tests + +```bash +python -m pytest tests/models/qwen3_next/test_qwen3_next_norm.py -v +python tools/validation/operators/check_operator.py --op qwen3_next_rms_norm --candidate cuda \ + --device cuda --dtype bf16 --check-grad +``` + +## Known Limitations + +- CUDA only; no ROCm, Ascend or Triton backend. +- 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 on [Gated RMSNorm](qwen3-next-rms-norm-gated.md). +- 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-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/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png new file mode 100644 index 000000000..eaf59c418 Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png new file mode 100644 index 000000000..5cd70b5cf Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json new file mode 100644 index 000000000..8a931b315 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json @@ -0,0 +1,400 @@ +{ + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": "822b085cd5b042661b2bda021414a2caa796bef0", + "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" + }, + "eps": 1e-06, + "dtype": "bfloat16", + "ops": { + "zero_centred_rmsnorm": { + "hidden": 2048, + "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": 25.200000032782555, + "backward_us": 556.6880106925964 + }, + "4096": { + "forward_us": 39.503999054431915, + "backward_us": 1260.4320049285889 + }, + "16384": { + "forward_us": 122.41600081324577, + "backward_us": 7879.663944244385 + }, + "65536": { + "forward_us": 427.5680035352707, + "backward_us": 30925.631523132324 + } + }, + "rl-kernel PyTorch reference": { + "1024": { + "forward_us": 71.45600020885468, + "backward_us": 595.3599810600281 + }, + "4096": { + "forward_us": 156.92799538373947, + "backward_us": 537.9360020160675 + }, + "16384": { + "forward_us": 578.4800052642822, + "backward_us": 992.5920069217682 + }, + "65536": { + "forward_us": 2150.12788772583, + "backward_us": 2983.504056930542 + } + }, + "transformers Qwen3NextRMSNorm": { + "1024": { + "forward_us": 70.25599852204323, + "backward_us": 599.3280112743378 + }, + "4096": { + "forward_us": 118.12799796462059, + "backward_us": 541.1999821662903 + }, + "16384": { + "forward_us": 409.2479944229126, + "backward_us": 1025.5680084228516 + }, + "65536": { + "forward_us": 1452.351987361908, + "backward_us": 3140.112042427063 + } + }, + "vLLM GemmaRMSNorm (forward only)": { + "1024": { + "forward_us": 67.4239993095398 + }, + "4096": { + "forward_us": 119.00799721479416 + }, + "16384": { + "forward_us": 407.3439985513687 + }, + "65536": { + "forward_us": 1451.9200325012207 + } + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "1024": { + "forward_us": 14.944000169634819 + }, + "4096": { + "forward_us": 17.008000053465366 + }, + "16384": { + "forward_us": 32.94399939477444 + }, + "65536": { + "forward_us": 90.43200314044952 + } + } + } + }, + "gated_rmsnorm": { + "hidden": 128, + "accuracy": { + "257": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0018348192997423702, + "dweight_max_abs_over_absmax": 0.0021439847223899285, + "dgate_max_abs_over_absmax": 0.0027467772330278155, + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0018348192997423702, + "dweight_max_abs_over_absmax": 0.0021439847223899285, + "dgate_max_abs_over_absmax": 0.0027467772330278155, + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "dx_max_abs_over_absmax": 0.004283006956510838, + "dweight_max_abs_over_absmax": 0.0053612828842400685, + "dgate_max_abs_over_absmax": 0.004464247864523204, + "forward_max_abs": 0.041932799726739134, + "forward_correctly_rounded_fraction": 0.6546084880828857 + }, + "vLLM RMSNormGated (forward only)": { + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + } + }, + "4096": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0028905074752616756, + "dweight_max_abs_over_absmax": 0.0017551573197679504, + "dgate_max_abs_over_absmax": 0.002440494893713147, + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0028905074752616756, + "dweight_max_abs_over_absmax": 0.0017551573197679504, + "dgate_max_abs_over_absmax": 0.002440494893713147, + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "dx_max_abs_over_absmax": 0.006106839188157589, + "dweight_max_abs_over_absmax": 0.002835271706901685, + "dgate_max_abs_over_absmax": 0.003959744684751284, + "forward_max_abs": 0.06659807537061546, + "forward_correctly_rounded_fraction": 0.6562480926513672 + }, + "vLLM RMSNormGated (forward only)": { + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999923706054688 + } + } + }, + "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": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "vLLM RMSNormGated (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + } + }, + "latency": { + "rl-kernel CUDA": { + "4096": { + "forward_us": 20.800000056624413, + "backward_us": 1269.7439789772034 + }, + "16384": { + "forward_us": 27.168000116944313, + "backward_us": 3861.135959625244 + }, + "65536": { + "forward_us": 64.00000303983688, + "backward_us": 14388.879776000977 + }, + "262144": { + "forward_us": 235.23200303316116, + "backward_us": 114274.65438842773 + } + }, + "rl-kernel PyTorch reference": { + "4096": { + "forward_us": 81.31200075149536, + "backward_us": 560.6080293655396 + }, + "16384": { + "forward_us": 89.72799777984619, + "backward_us": 550.2559840679169 + }, + "65536": { + "forward_us": 209.3760073184967, + "backward_us": 573.7600028514862 + }, + "262144": { + "forward_us": 774.8000025749207, + "backward_us": 1376.08003616333 + } + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "4096": { + "forward_us": 90.01599997282028, + "backward_us": 577.888011932373 + }, + "16384": { + "forward_us": 91.80799871683121, + "backward_us": 601.1359989643097 + }, + "65536": { + "forward_us": 196.46400213241577, + "backward_us": 623.5039830207825 + }, + "262144": { + "forward_us": 697.2479820251465, + "backward_us": 1528.384029865265 + } + }, + "vLLM RMSNormGated (forward only)": { + "4096": { + "forward_us": 49.775999039411545 + }, + "16384": { + "forward_us": 51.024001091718674 + }, + "65536": { + "forward_us": 50.20799860358238 + }, + "262144": { + "forward_us": 81.4880020916462 + } + } + } + } + } +} diff --git a/rl_engine/_C.pyi b/rl_engine/_C.pyi index fb3c4379c..c0f9cb02c 100644 --- a/rl_engine/_C.pyi +++ b/rl_engine/_C.pyi @@ -256,16 +256,38 @@ 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_gated_forward( + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + eps: float, + weight_offset: float = ..., + activation: int = ..., +) -> list[torch.Tensor]: ... +def rmsnorm_gated_backward_dx( + dy: torch.Tensor, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + rstd: torch.Tensor, + weight_offset: float = ..., + activation: int = ..., ) -> 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..60cdbffd3 100644 --- a/rl_engine/backends/cuda/norm/rmsnorm.py +++ b/rl_engine/backends/cuda/norm/rmsnorm.py @@ -4,9 +4,21 @@ 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 _fold_dweight_rows(rows: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: + """The single left-fold entrypoint for this backend's dweight reductions. + + Both the plain and the gated backward route through here so the file keeps + one auditable reduction path; the ascending-row fp32 fold is what makes + dweight independent of the batch layout. + """ + return reduce_rows_fp32(rows).to(dtype) + + +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 +27,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 +50,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 +74,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 +83,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 +101,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. @@ -89,7 +109,7 @@ def backward(ctx, grad_out): # Multiplication is part of the pre-existing mask contract, including # IEEE propagation for non-finite inactive contributions. rows = rows * mask.to(dtype=rows.dtype).unsqueeze(-1) - dw = reduce_rows_fp32(rows).to(weight.dtype) + dw = _fold_dweight_rows(rows, weight.dtype) record_backward( "rms_norm", kernel_id=( @@ -101,16 +121,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 +141,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 +153,236 @@ 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 + + +# --------------------------------------------------------------------------- # +# Gated RMSNorm (Qwen3-Next GDN block) +# --------------------------------------------------------------------------- # + + +def _require_cuda_symbols(what: str, *names: str) -> None: + """Raise when the compiled kernels backing ``what`` are missing. + + The registry treats a backend whose construction raises as unavailable and + falls through, so calling this from ``__init__`` is what lets a CUDA-first + priority list degrade to the PyTorch reference on a build without the + extension. Mirrors ``_require_cuda_activation`` in the activation ops. + """ + if not _EXT_AVAILABLE or _C is None: + raise RuntimeError(f"{what} requires the compiled rl_engine._C extension.") + 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. " + "Rebuild the extension with csrc/cuda/norm/rmsnorm.cu." + ) + + +#: Gate activations understood by the CUDA kernel, in binding order. ``swish`` is +#: an alias for ``silu``, as in vLLM's GDN block, which maps ``output_gate_type`` +#: "swish" to "silu" before constructing ``RMSNormGated``. +_GATE_ACTIVATIONS = {"silu": 0, "swish": 0, "sigmoid": 1} + + +def _check_gate_activation(activation: str) -> int: + if activation not in _GATE_ACTIVATIONS: + raise ValueError( + f"activation must be one of {sorted(_GATE_ACTIVATIONS)}, got {activation!r}" + ) + return _GATE_ACTIVATIONS[activation] + + +def _gate_activation_fp32(gate: torch.Tensor, activation: int) -> torch.Tensor: + """act(gate) in fp32, matching the kernel's ``gate_activation``.""" + gate32 = gate.float() + return torch.nn.functional.silu(gate32) if activation == 0 else torch.sigmoid(gate32) + + +def _gate_activation_grad_fp32(gate: torch.Tensor, activation: int) -> torch.Tensor: + """d act(gate) / d gate in fp32, matching ``gate_activation_grad``.""" + gate32 = gate.float() + sigma = torch.sigmoid(gate32) + if activation == 0: + return sigma * (1.0 + gate32 * (1.0 - sigma)) + return sigma * (1.0 - sigma) + + +class RMSNormGatedCuda(torch.autograd.Function): + """Autograd wrapper for the gated CUDA RMSNorm. + + Forward is the fused kernel. Backward is assembled from deterministic + pieces: ``dx`` from a row-local CUDA kernel, ``dweight`` from fp32 row + contributions reduced by the ascending-row left fold, and ``dgate`` purely + elementwise in fp32 (no reduction, so batch invariance is trivial). + """ + + @staticmethod + def forward(ctx, x, weight, gate, eps=1e-6, weight_offset=0.0, activation=0): + """ + Forward: + y = x * rsqrt(mean(x^2) + eps) * (weight_offset + weight) * act(gate) + + Input: + x, gate: [T, H], fp16/bf16/fp32 CUDA tensors of matching dtype + weight: [H] + activation: 0 = silu/swish, 1 = sigmoid + """ + assert x.is_cuda and weight.is_cuda and gate.is_cuda, "inputs must be CUDA tensors" + assert x.is_contiguous() and weight.is_contiguous() and gate.is_contiguous() + assert x.dim() == 2, "x must be [T, H]" + assert weight.dim() == 1, "weight must be [H]" + assert gate.shape == x.shape, "gate must match x" + assert _EXT_AVAILABLE and hasattr( + _C, "rmsnorm_gated_forward" + ), "Gated RMSNorm CUDA extension is unavailable. Rebuild with csrc/cuda/norm/rmsnorm.cu." + + y, rstd = _C.rmsnorm_gated_forward( + x, weight, gate, float(eps), float(weight_offset), int(activation) + ) + + ctx.save_for_backward(x, weight, gate, rstd) + ctx.eps = eps + ctx.weight_offset = float(weight_offset) + ctx.activation = int(activation) + + return y + + @staticmethod + def backward(ctx, grad_out): + x, weight, gate, rstd = ctx.saved_tensors + dy = grad_out.contiguous() + act = ctx.activation + + dx = _C.rmsnorm_gated_backward_dx(dy, x, weight, gate, rstd, ctx.weight_offset, act) + + # dweight: the gate is a per-element constant here, so the ungated row + # contributions apply once dy carries act(gate). + gate_act = _gate_activation_fp32(gate, act) + rows = rmsnorm_dweight_rows_fp32(x, dy.float() * gate_act, rstd=rstd) + dw = _fold_dweight_rows(rows, weight.dtype) + + # dgate: row-local and reduction-free. + normed = x.float() * rstd.unsqueeze(-1) + # Same guard as the kernels: an unconditional `+ 0.0` turns -0.0 weights into +0.0. + scale = weight.float() + if ctx.weight_offset != 0.0: + scale = scale + ctx.weight_offset + dgate = (dy.float() * normed * scale * _gate_activation_grad_fp32(gate, act)).to(gate.dtype) + + record_backward( + "rms_norm_gated", + kernel_id=( + "rl_engine._C.rmsnorm_gated_backward_dx" + "+rl_engine.ops.autograd.vjp_fp32.rmsnorm_dweight_rows_fp32" + "+rl_engine.ops.autograd.vjp_fp32.reduce_rows_fp32" + ), + impl="cuda_rmsnorm_gated_dx_declared_fp32_rowfold_dw", + family="cuda", + ) + + return dx, dw, dgate, None, None, None + + +def rmsnorm_gated_cuda(x, weight, gate, eps=1e-6, weight_offset=0.0, activation="silu"): + """ + use: + y = rmsnorm_gated_cuda(x, weight, gate) + y = rmsnorm_gated_cuda(x, weight, gate, activation="sigmoid") + """ + act = _check_gate_activation(activation) + return RMSNormGatedCuda.apply(x, weight, gate, eps, weight_offset, act) + + +class Qwen3NextRMSNormGatedCudaOp: + """CUDA gated RMSNorm for the Qwen3-Next GDN block. + + Deliberately not a subclass of :class:`RMSNormCudaOp`: it takes an extra + required tensor, so it cannot stand in for one. + + ``out = x * rstd * weight * silu(gate)``, every multiply in fp32 with a + single cast at the store. The weight is plain, not zero-centred, matching + vLLM's ``RMSNormGated`` with ``norm_before_gate=True`` and ``group_size=None`` + -- which is how the GDN block constructs it. The op has no ``norm_before_gate`` + or ``group_size`` parameter, so other configurations are not implemented. + + ``activation`` is fixed at construction; the registry constructs the + released config's ``"silu"``. + """ + + backward_impl = "cuda_rmsnorm_gated_dx_declared_fp32_rowfold_dw" + + #: The gated weight is plain; kept as an attribute so the surface matches + #: the ungated op and a zero-centred variant stays one subclass away. + weight_offset = 0.0 + + def __init__(self, activation: str = "silu") -> None: + _check_gate_activation(activation) + self.activation = activation + _require_cuda_symbols( + "Gated CUDA RMSNorm", "rmsnorm_gated_forward", "rmsnorm_gated_backward_dx" + ) + + def __call__(self, x, weight, gate, *, eps=1e-6): + return self.forward(x, weight, gate, eps=eps) + + def forward(self, x, weight, gate, *, eps=1e-6): + if gate.shape != x.shape: + raise ValueError(f"gate must match x, got {tuple(gate.shape)} vs {tuple(x.shape)}") + hidden = x.shape[-1] + x_2d = x.contiguous().view(-1, hidden) + gate_2d = gate.contiguous().view(-1, hidden) + y_2d = rmsnorm_gated_cuda( + x_2d, + weight.contiguous(), + gate_2d, + eps=eps, + weight_offset=self.weight_offset, + activation=self.activation, + ) + return y_2d.view_as(x) + + def parameter_vjp_contributions_fp32(self, *, x, weight, gate, grad_output, eps=1e-6): + if gate.shape != x.shape: + raise ValueError(f"gate must match x, got {tuple(gate.shape)} vs {tuple(x.shape)}") + hidden = x.shape[-1] + act = _GATE_ACTIVATIONS[self.activation] + _, rstd = _C.rmsnorm_gated_forward( + x.contiguous().reshape(-1, hidden), + weight.contiguous(), + gate.contiguous().reshape(-1, hidden), + float(eps), + float(self.weight_offset), + act, + ) + rows = rmsnorm_dweight_rows_fp32( + x, + grad_output.float() * _gate_activation_fp32(gate, act), + rstd=rstd.reshape(x.shape[:-1]), + ) return {"weight": rows} diff --git a/rl_engine/config/workload.py b/rl_engine/config/workload.py index 69b21c20f..79f93208a 100644 --- a/rl_engine/config/workload.py +++ b/rl_engine/config/workload.py @@ -264,10 +264,30 @@ def load_manifest(path: str | Path | None = None) -> WS1Manifest: raw = json.load(fh) if not isinstance(raw, dict): raise WorkloadError("manifest root must be a JSON object") - validate_manifest(raw) + if raw.get("scope") == "qwen3_next_norm_operators": + from rl_engine.validation.models.qwen3_next_workload import validate_norm_manifest + + validate_norm_manifest(raw) + else: + validate_manifest(raw) return WS1Manifest(raw=raw, path=manifest_path) +def workload_report(manifest: WS1Manifest) -> dict[str, Any]: + """The ``workload`` block the C3/C4 gate scripts attach to their JSON report. + + ``full_model_evidence`` is whatever the manifest declares, and ``None`` when + it declares nothing. + """ + return { + "workload_id": manifest.workload_id, + "scope": manifest.raw.get("scope", "qwen3_8b_dense"), + "fixture_identity_sha256": manifest.raw["fixture_identity_sha256"], + "model_id": manifest.model_identity["model_id"], + "full_model_evidence": manifest.raw.get("full_model_evidence"), + } + + def validate_manifest(raw: Mapping[str, Any]) -> None: """Hard-fail if any required C2 pin is missing or inconsistent.""" missing = [k for k in _REQUIRED_TOP_LEVEL if k not in raw] @@ -292,20 +312,25 @@ def validate_manifest(raw: Mapping[str, Any]) -> None: ) -def _validate_model_identity(identity: Mapping[str, Any]) -> None: +def _validate_model_identity( + identity: Mapping[str, Any], + *, + fingerprint: Mapping[str, Any] = _OFFICIAL_FINGERPRINT, + model_label: str = "Qwen3-8B Dense", +) -> None: for key in ("model_id", "revision", "config_fingerprint", "weight_snapshot"): if key not in identity: raise WorkloadError(f"model_identity missing {key!r}") fp = identity["config_fingerprint"] if not isinstance(fp, Mapping): raise WorkloadError("config_fingerprint must be an object") - for key, expected in _OFFICIAL_FINGERPRINT.items(): + for key, expected in fingerprint.items(): if key not in fp: raise WorkloadError(f"config_fingerprint missing {key!r}") if fp[key] != expected: raise WorkloadError( f"config_fingerprint {key}={fp[key]!r} does not match official " - f"Qwen3-8B Dense pin {expected!r}; architecture shrink is forbidden" + f"{model_label} pin {expected!r}; architecture shrink is forbidden" ) if not identity.get("exit_forbids_architecture_shrink", False): raise WorkloadError("exit_forbids_architecture_shrink must be true") diff --git a/rl_engine/reference/linear_attn/__init__.py b/rl_engine/reference/linear_attn/__init__.py new file mode 100644 index 000000000..80a2a86b9 --- /dev/null +++ b/rl_engine/reference/linear_attn/__init__.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Linear-attention operators (Gated DeltaNet and friends).""" + +from rl_engine.reference.linear_attn.causal_conv1d import CausalConv1dUpdateOp +from rl_engine.reference.linear_attn.gated_delta_rule import ( + NULL_BLOCK_ID, + SOFTPLUS_THRESHOLD, + GatedDeltaRuleRecurrentStepOp, +) + +__all__ = [ + "CausalConv1dUpdateOp", + "GatedDeltaRuleRecurrentStepOp", + "NULL_BLOCK_ID", + "SOFTPLUS_THRESHOLD", +] diff --git a/rl_engine/reference/linear_attn/causal_conv1d.py b/rl_engine/reference/linear_attn/causal_conv1d.py new file mode 100644 index 000000000..97886c8c1 --- /dev/null +++ b/rl_engine/reference/linear_attn/causal_conv1d.py @@ -0,0 +1,184 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Causal depthwise conv1d single-token state update (WS1 ground truth, RFC #428 C6). + +The trainer-side reference for ``causal_conv1d_update``, which vLLM calls once +per decode token before the Gated DeltaNet recurrence. + +The incoming token is rounded to the cache dtype. Products are rounded in +operand dtype before sequential FP32 accumulation, starting from bias (or zero). +This matters for BF16 caches: promoting operands before multiplication loses the +provider's product rounding. Activation may still differ at transcendental ULP +scale; this reference does not establish model-level bitwise equality. + +""" + +from __future__ import annotations + +import torch + +from rl_engine.reference.linear_attn.gated_delta_rule import NULL_BLOCK_ID, _validate_state_indices + +__all__ = ["CausalConv1dUpdateOp"] + + +class CausalConv1dUpdateOp: + """One decode token through the paged causal-conv1d cache. + + ============== ========================== ========================== + tensor shape notes + ============== ========================== ========================== + ``x`` ``[B, dim]`` + ``conv_state`` ``[num_blocks, dim, W-1]`` paged; ``dim_first`` layout + ``weight`` ``[dim, W]`` + ``bias`` ``[dim]`` or ``None`` + ``indices`` ``[B]`` ``<= 0`` skips (vLLM: ``== 0``) + ============== ========================== ========================== + + Returns ``(out, conv_state)``; the state is updated out of place. + """ + + def __call__( + self, + x: torch.Tensor, + conv_state: torch.Tensor, + weight: torch.Tensor, + conv_state_indices: torch.Tensor, + *, + bias: torch.Tensor | None = None, + activation: str | None = "silu", + dim_first: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + return self.forward( + x, + conv_state, + weight, + conv_state_indices, + bias=bias, + activation=activation, + dim_first=dim_first, + ) + + def forward( + self, + x: torch.Tensor, + conv_state: torch.Tensor, + weight: torch.Tensor, + conv_state_indices: torch.Tensor, + *, + bias: torch.Tensor | None = None, + activation: str | None = "silu", + dim_first: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + return self._update( + x, + conv_state, + weight, + conv_state_indices, + bias=bias, + activation=activation, + dim_first=dim_first, + output_dtype=x.dtype, + ) + + def forward_fp32( + self, + x: torch.Tensor, + conv_state: torch.Tensor, + weight: torch.Tensor, + conv_state_indices: torch.Tensor, + *, + bias: torch.Tensor | None = None, + activation: str | None = "silu", + dim_first: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Ground truth: fp32 output, so only the cache dtype rounds.""" + return self._update( + x, + conv_state, + weight, + conv_state_indices, + bias=bias, + activation=activation, + dim_first=dim_first, + output_dtype=torch.float32, + ) + + @staticmethod + def _update( + x, + conv_state, + weight, + conv_state_indices, + *, + bias, + activation, + dim_first, + output_dtype, + ): + if activation not in (None, "silu", "swish"): + raise ValueError(f"activation must be None, 'silu' or 'swish', got {activation!r}") + if x.dim() != 2: + raise ValueError(f"x must be 2-D [B, dim], got {tuple(x.shape)}") + if conv_state.dim() != 3: + raise ValueError( + f"conv_state must be 3-D [num_blocks, dim, W-1], got {tuple(conv_state.shape)}" + ) + if conv_state_indices.dim() != 1 or conv_state_indices.shape[0] != x.shape[0]: + raise ValueError("conv_state_indices must be 1-D with one entry per row of x") + + # The "SD" layout stores (state_len, dim); the kernels want (dim, state_len). + state = conv_state if dim_first else conv_state.transpose(-1, -2) + dim, tail = state.shape[-2], state.shape[-1] + if weight.ndim != 2 or dim <= 0 or weight.shape[-1] <= 0: + raise ValueError("weight must be 2-D [dim, W] with positive dimensions") + if x.shape[1] != dim: + raise ValueError("x and conv_state must have the same dim") + _validate_state_indices(conv_state_indices, x.shape[0], state.shape[0], x.device) + for name, tensor in ( + ("x", x), + ("conv_state", conv_state), + ("weight", weight), + ("bias", bias), + ): + if tensor is not None and (tensor.device != x.device or not tensor.is_floating_point()): + raise ValueError(f"{name} must be floating point on the input device") + if bias is not None and bias.shape != (dim,): + raise ValueError("bias must have shape [dim]") + width = weight.shape[-1] + if weight.shape[0] != dim: + raise ValueError(f"weight must be [dim, W] with dim={dim}, got {tuple(weight.shape)}") + if tail != width - 1: + raise ValueError(f"conv_state tail must be W-1={width - 1}, got {tail}") + + active = conv_state_indices > NULL_BLOCK_ID + out = torch.zeros(x.shape, dtype=torch.float32, device=x.device) + new_state = conv_state.clone() # the cache dtype is the contract here + if not bool(active.any()): + return out.to(output_dtype), new_state + + rows = torch.nonzero(active, as_tuple=False).flatten() + blocks = conv_state_indices[rows].long() + + # The incoming token is rounded to the cache dtype before it is used. + token = x[rows].to(conv_state.dtype) + window = torch.cat([state[blocks], token.unsqueeze(-1)], dim=-1) # [R, dim, W] + + # Match the provider's product rounding and bias-before-taps order. + acc = torch.zeros(len(rows), dim, dtype=torch.float32, device=x.device) + if bias is not None: + acc = acc + bias.float() + for tap in range(width): + product = window[..., tap] * weight[:, tap].unsqueeze(0) + acc = acc + product.float() + if activation in ("silu", "swish"): + acc = acc / (1.0 + torch.exp(-acc)) + + out[rows] = acc + rolled = window[..., 1:] + if dim_first: + new_state[blocks] = rolled.to(conv_state.dtype) + else: + new_state[blocks] = rolled.transpose(-1, -2).to(conv_state.dtype) + return out.to(output_dtype), new_state diff --git a/rl_engine/reference/linear_attn/gated_delta_rule.py b/rl_engine/reference/linear_attn/gated_delta_rule.py new file mode 100644 index 000000000..442f1f6ac --- /dev/null +++ b/rl_engine/reference/linear_attn/gated_delta_rule.py @@ -0,0 +1,334 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Gated DeltaNet single-token recurrent step (WS1 ground truth for RFC #428 C6). + +This is the trainer-side reference for what vLLM runs during rollout decode: +``fused_recurrent_gated_delta_rule_packed_decode``. That kernel -- not +``fused_sigmoid_gating_delta_rule_update`` -- is the path a pure RL rollout +takes, because ``VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE`` defaults to true and +a decode-only, non-speculative batch returns early into it. + +Three things are transcribed from the kernel rather than from the HuggingFace +model, because they differ: + +* **The gating is fused.** ``beta = sigmoid(b)`` and + ``g = -exp(A_log) * softplus(a + dt_bias)`` are computed here, in fp32, with + the kernel's ``softplus`` threshold branch. HF computes them separately in + PyTorch, which is a different rounding path. +* **No ``repeat_interleave``.** The kernel indexes ``i_h = i_hv // (HV // H)``, + so q/k stay at ``H`` heads while v has ``HV``. HF materializes the repeat. +* **The QK norm is an L2 norm over a plain sum**, ``x / sqrt(sum(x*x) + 1e-6)``, + not an RMSNorm and not ``F.normalize``. It divides by ``sqrt`` rather than + multiplying by ``rsqrt``; the two differ in the last bit, so the division is + kept. + +``scale`` is applied to ``q`` *after* the L2 norm; ``k`` is never scaled. + +State ABI (mirrored, not reinvented) +------------------------------------ +The recurrent state is paged: ``[num_blocks, HV, V, K]``, addressed per +sequence by ``ssm_state_indices``. Index ``<= 0`` is ``NULL_BLOCK_ID`` and means +"skip": the kernel writes zeros to the output and leaves the block untouched. +The accumulator is fp32 throughout the step, and the store rounds to the state +tensor's dtype -- which upstream allows to be fp32 *or* bf16. That rounding is +part of the recurrence and compounds across tokens, so it is modelled here +rather than skipped. +""" + +from __future__ import annotations + +import torch + +__all__ = ["GatedDeltaRuleRecurrentStepOp", "NULL_BLOCK_ID", "SOFTPLUS_THRESHOLD"] + +#: Paged-state sentinel. Both goldens skip a row whose index is ``<= NULL_BLOCK_ID``: +#: that is the recurrent provider's semantics (``state_idx <= 0``). vLLM's conv +#: provider skips only ``== null_block_id`` (0), so a negative index is a real index +#: there; see the design note's NULL_BLOCK_ID table. +NULL_BLOCK_ID = 0 + +#: Above this, softplus is the identity (matches the kernel's constexpr). +SOFTPLUS_THRESHOLD = 20.0 + +#: Width of the fixed-order reduction chunks. The contraction order must not +#: depend on the batch layout, so it is pinned here exactly as +#: :func:`~rl_engine.reference.norm.rms_norm.shape_invariant_rstd` +#: pins the RMSNorm statistic. +_REDUCTION_CHUNK = 32 + + +def _validate_state_indices(indices, batch, blocks, device): + """Each active row owns one cache block; inactive sentinels may repeat.""" + if indices.ndim != 1 or indices.shape[0] != batch: + raise ValueError("state indices must be 1-D with one entry per sequence") + if indices.dtype not in (torch.int32, torch.int64): + raise ValueError("state indices must have int32 or int64 dtype") + if indices.device != device: + raise ValueError("state indices must be on the input device") + active = indices[indices > NULL_BLOCK_ID] + if bool((active >= blocks).any()): + raise ValueError("active state index is out of range") + if active.unique().numel() != active.numel(): + raise ValueError("active state indices must be unique") + + +def _chunked_sum(x: torch.Tensor) -> torch.Tensor: + """Sum the last dim in a fixed 32-wide chunk order. + + The single reduction primitive for this module. Deliberately not + ``sum``/``matmul``/``einsum`` on the whole axis: their reduction order is + unspecified and may vary with shape, which is what a batch-invariant claim + cannot tolerate. Mirrors + :func:`~rl_engine.kernels.ops.pytorch.norm.rms_norm.shape_invariant_rstd`. + """ + tail = x.shape[-1] + if tail % _REDUCTION_CHUNK != 0: + # Zero-pad to the next chunk boundary instead of falling back to an + # unpinned ``sum``: exact zeros add nothing, so the result is the same + # fixed chunk order for every width. + x = torch.nn.functional.pad(x, (0, _REDUCTION_CHUNK - tail % _REDUCTION_CHUNK)) + tail = x.shape[-1] + return ( + x.reshape(*x.shape[:-1], tail // _REDUCTION_CHUNK, _REDUCTION_CHUNK).sum(dim=-1).sum(dim=-1) + ) + + +def _fixed_order_contract(mat: torch.Tensor, vec: torch.Tensor) -> torch.Tensor: + """``sum(mat * vec, dim=-1)`` in a fixed chunk order.""" + return _chunked_sum(mat * vec) + + +def _l2_normalize(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor: + """``x / sqrt(sum(x*x) + eps)`` over the last dim, in fp32. + + A plain sum, not a mean: this is an L2 norm, not an RMSNorm. The division is + kept rather than folded into a reciprocal-sqrt multiply because the kernel + divides. + """ + return x / torch.sqrt(_chunked_sum(x * x) + eps).unsqueeze(-1) + + +def _softplus(x: torch.Tensor, threshold: float = SOFTPLUS_THRESHOLD) -> torch.Tensor: + """The kernel's branched softplus; the branch matters for large ``a + dt_bias``.""" + small = x <= threshold + safe = torch.where(small, x, torch.zeros_like(x)) + return torch.where(small, torch.log1p(torch.exp(safe)), x) + + +class GatedDeltaRuleRecurrentStepOp: + """One decode token of the Gated DeltaNet recurrence, over a paged state. + + Shapes follow the provider, not the model definition: + + ============= ====================================== =================== + tensor shape notes + ============= ====================================== =================== + ``mixed_qkv`` ``[B, H*K + H*K + HV*V]`` q | k | v, packed + ``a``, ``b`` ``[B, HV]`` + ``A_log`` ``[HV]`` fp32 + ``dt_bias`` ``[HV]`` fp32 + ``state`` ``[num_blocks, HV, V, K]`` paged, V-major + ``indices`` ``[B]`` ``<= 0`` skips + ============= ====================================== =================== + + Returns ``(out, state)`` with ``out`` of shape ``[B, 1, HV, V]``. The state + is updated out of place; pass the result back to continue the recurrence. + """ + + def __call__( + self, + mixed_qkv: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + state: torch.Tensor, + ssm_state_indices: torch.Tensor, + *, + scale: float, + num_k_heads: int, + use_qk_l2norm: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + return self.forward( + mixed_qkv, + a, + b, + A_log, + dt_bias, + state, + ssm_state_indices, + scale=scale, + num_k_heads=num_k_heads, + use_qk_l2norm=use_qk_l2norm, + ) + + def forward( + self, + mixed_qkv: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + state: torch.Tensor, + ssm_state_indices: torch.Tensor, + *, + scale: float, + num_k_heads: int, + use_qk_l2norm: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Step once, rounding the stored state to ``state.dtype``.""" + return self._step( + mixed_qkv, + a, + b, + A_log, + dt_bias, + state, + ssm_state_indices, + scale=scale, + num_k_heads=num_k_heads, + use_qk_l2norm=use_qk_l2norm, + state_dtype=state.dtype, + output_dtype=mixed_qkv.dtype, + ) + + def forward_fp32( + self, + mixed_qkv: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + state: torch.Tensor, + ssm_state_indices: torch.Tensor, + *, + scale: float, + num_k_heads: int, + use_qk_l2norm: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Ground truth: keep the state and the output in fp32. + + The difference from :meth:`forward` is the per-token state rounding, so + running both and comparing isolates how much that rounding costs. + """ + return self._step( + mixed_qkv, + a, + b, + A_log, + dt_bias, + state, + ssm_state_indices, + scale=scale, + num_k_heads=num_k_heads, + use_qk_l2norm=use_qk_l2norm, + state_dtype=torch.float32, + output_dtype=torch.float32, + ) + + @staticmethod + def _step( + mixed_qkv, + a, + b, + A_log, + dt_bias, + state, + ssm_state_indices, + *, + scale, + num_k_heads, + use_qk_l2norm, + state_dtype, + output_dtype, + ): + if mixed_qkv.dim() != 2: + raise ValueError(f"mixed_qkv must be 2-D [B, D], got {tuple(mixed_qkv.shape)}") + if state.dim() != 4: + raise ValueError(f"state must be 4-D [num_blocks, HV, V, K], got {tuple(state.shape)}") + + batch = mixed_qkv.shape[0] + hv, v_dim, k_dim = state.shape[-3:] + if isinstance(num_k_heads, bool) or not isinstance(num_k_heads, int) or num_k_heads <= 0: + raise ValueError("num_k_heads must be a positive integer") + if min(hv, v_dim, k_dim) <= 0: + raise ValueError("state head dimensions must be positive") + heads = num_k_heads + _validate_state_indices(ssm_state_indices, batch, state.shape[0], mixed_qkv.device) + for name, tensor in ( + ("a", a), + ("b", b), + ("A_log", A_log), + ("dt_bias", dt_bias), + ("state", state), + ): + if tensor.device != mixed_qkv.device: + raise ValueError(f"{name} must be on the input device") + if not tensor.is_floating_point(): + raise ValueError(f"{name} must be floating point") + if not mixed_qkv.is_floating_point(): + raise ValueError("mixed_qkv must be floating point") + if A_log.shape != (hv,) or dt_bias.shape != (hv,): + raise ValueError("A_log and dt_bias must have shape [HV]") + if hv % heads != 0: + raise ValueError(f"HV={hv} must be a multiple of num_k_heads={heads}") + if a.shape != (batch, hv) or b.shape != (batch, hv): + raise ValueError( + f"a/b must be [B, HV] = {(batch, hv)}, got {tuple(a.shape)} / {tuple(b.shape)}" + ) + expected = heads * k_dim * 2 + hv * v_dim + if mixed_qkv.shape[1] != expected: + raise ValueError( + f"mixed_qkv last dim must be {expected} (q|k|v packed), " + f"got {mixed_qkv.shape[1]}" + ) + + group = hv // heads + qkv32 = mixed_qkv.float() + + # Unpack q | k | v. q and k carry H heads, v carries HV. + q = qkv32[:, : heads * k_dim].reshape(batch, heads, k_dim) + k = qkv32[:, heads * k_dim : 2 * heads * k_dim].reshape(batch, heads, k_dim) + v = qkv32[:, 2 * heads * k_dim :].reshape(batch, hv, v_dim) + + if use_qk_l2norm: + q = _l2_normalize(q) + k = _l2_normalize(k) + q = q * scale + + # i_h = i_hv // (HV // H): index, do not materialize a repeat. + head_of = torch.arange(hv, device=q.device) // group + q = q[:, head_of, :] # [B, HV, K] + k = k[:, head_of, :] + + # Fused gating, fp32, with the kernel's softplus branch. + decay = -torch.exp(A_log.float()) * _softplus(a.float() + dt_bias.float()) + beta = torch.sigmoid(b.float()) + + active = ssm_state_indices > NULL_BLOCK_ID + out = torch.zeros(batch, 1, hv, v_dim, dtype=torch.float32, device=q.device) + # The returned state carries `state_dtype`, not the caller's: forward_fp32 + # exists precisely to run the recurrence without the per-token rounding, + # so it must be able to widen a bf16 cache to fp32. + new_state = state.to(state_dtype).clone() + if not bool(active.any()): + return out.to(output_dtype), new_state + + rows = torch.nonzero(active, as_tuple=False).flatten() + blocks = ssm_state_indices[rows].long() + + h = state[blocks].float() # [R, HV, V, K] + h = h * torch.exp(decay[rows]).unsqueeze(-1).unsqueeze(-1) + + k_sel, q_sel, v_sel = k[rows], q[rows], v[rows] + # v -= h @ k, then v *= beta, then h += outer(v, k), then o = h @ q. + v_sel = v_sel - _fixed_order_contract(h, k_sel.unsqueeze(-2)) + v_sel = v_sel * beta[rows].unsqueeze(-1) + h = h + v_sel.unsqueeze(-1) * k_sel.unsqueeze(-2) + o = _fixed_order_contract(h, q_sel.unsqueeze(-2)) + + out[rows, 0] = o + # The store rounds; that rounding is part of the recurrence. + new_state[blocks] = h.to(state_dtype) + return out.to(output_dtype), new_state 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..14ecdcfa6 --- /dev/null +++ b/rl_engine/reference/norm/qwen3_next_rms_norm.py @@ -0,0 +1,161 @@ +# 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`` and, for the gated +pair, ``docs/operators/qwen3-next-rms-norm-gated.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 ``docs/operators/qwen3-next-rms-norm-gated.md``. + + 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)}" + ) + if gate.dtype != x.dtype or gate.device != x.device: + raise ValueError("gate must have the same dtype and device as x") + 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/rl_engine/runtime/registry.py b/rl_engine/runtime/registry.py index a190bbca1..9a6494d5d 100644 --- a/rl_engine/runtime/registry.py +++ b/rl_engine/runtime/registry.py @@ -159,6 +159,18 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): TRITON_RMS_NORM = "rl_engine.kernels.ops.triton.rmsnorm_triton.RMSNormTritonOp" PYTORCH_NATIVE_RMS_NORM = "rl_engine.kernels.ops.pytorch.norm.rms_norm.NativeRMSNormOp" + # Zero-centred RMSNorm (Qwen3-Next decoder / final norm) + CUDA_QWEN3_NEXT_RMS_NORM = "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp" + PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM = ( + "rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp" + ) + + # Gated RMSNorm (Qwen3-Next Gated DeltaNet block) + CUDA_RMS_NORM_GATED = "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp" + PYTORCH_NATIVE_RMS_NORM_GATED = ( + "rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp" + ) + # Generic fallback TRITON_GENERIC = "rl_engine.kernels.ops.triton.generic.TritonOp" PYTORCH_ATTN = "rl_engine.kernels.ops.pytorch.attention.NativeAttentionOp" @@ -594,6 +606,14 @@ def __init__(self): OpBackend.CUDA_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ], + "rms_norm_gated": [ + OpBackend.CUDA_RMS_NORM_GATED, + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + ], + "qwen3_next_rms_norm": [ + OpBackend.CUDA_QWEN3_NEXT_RMS_NORM, + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + ], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [ @@ -649,6 +669,8 @@ def __init__(self): ], "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.TRITON_SILU, OpBackend.PYTORCH_NATIVE_SILU], @@ -694,6 +716,8 @@ def __init__(self): OpBackend.TRITON_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [ OpBackend.TRITON_EMBEDDING, @@ -722,6 +746,8 @@ def __init__(self): "batch_invariant_logp": [OpBackend.PYTORCH_BATCH_INVARIANT_LOGP], "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.PYTORCH_NATIVE_SILU], @@ -762,6 +788,12 @@ def __init__(self): OpBackend.ASCEND_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ] + self._priority_map["npu"]["rms_norm_gated"] = [ + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + ] + self._priority_map["npu"]["qwen3_next_rms_norm"] = [ + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + ] self._priority_map["npu"]["embedding"] = [ OpBackend.ASCEND_EMBEDDING, OpBackend.PYTORCH_NATIVE_EMBEDDING, diff --git a/rl_engine/validation/models/qwen3_next_norm_manifest.json b/rl_engine/validation/models/qwen3_next_norm_manifest.json new file mode 100644 index 000000000..9cf5977ef --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_norm_manifest.json @@ -0,0 +1,759 @@ +{ + "version": "1.0", + "workload_id": "qwen3-next-80b-a3b-norm-c3-c4-v1", + "seed": 20260812, + "model_identity": { + "model_id": "Qwen/Qwen3-Next-80B-A3B-Instruct", + "revision": "9c7f2fbe84465e40164a94cc16cd30b6999b0cc7", + "config_fingerprint": { + "num_hidden_layers": 48, + "hidden_size": 2048, + "intermediate_size": 5120, + "num_attention_heads": 16, + "num_key_value_heads": 2, + "head_dim": 256, + "vocab_size": 151936, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "num_experts": 512, + "num_experts_per_tok": 10 + }, + "exit_forbids_architecture_shrink": true, + "weight_snapshot": { + "pin_method": "revision-and-sha256", + "total_size_bytes": 162649725440, + "index_file": "model.safetensors.index.json", + "index_sha256": "0e5274f538d750abdd7fd55a36bad64100f9e881d59c65f6a59f364d2f8b6dd7", + "content_hash_algorithm": "sha256-of-sorted-shard-records-v1", + "content_hash": "00a6a3b5573ccbac2e01fadb673c2cd7bca172d40ddd73f95ed7afa5cbdb9dcd", + "shards": [ + { + "filename": "model-00001-of-00041.safetensors", + "size_bytes": 3999619256, + "sha256": "d8908fd48650169854b5ed815cb01bfc1741152cc58795ec99188862c6a29e11" + }, + { + "filename": "model-00002-of-00041.safetensors", + "size_bytes": 3999841784, + "sha256": "c753c9bfaca220781d4030c3a99e69b4a256434c9d6ec223f5147edc265289df" + }, + { + "filename": "model-00003-of-00041.safetensors", + "size_bytes": 3999515584, + "sha256": "51aaa14dd50c5ab90c363227bfb1ac51182118f588e0141dcaadda012548407c" + }, + { + "filename": "model-00004-of-00041.safetensors", + "size_bytes": 3999842000, + "sha256": "82a33096134fb6e7751a423f142e79b1bbe89f45242b75397c2d7170fffa75bd" + }, + { + "filename": "model-00005-of-00041.safetensors", + "size_bytes": 3999842208, + "sha256": "41abda7bf93f27c36cb28382114f5b33defa6dbfd154abb6886a6fec20f9e479" + }, + { + "filename": "model-00006-of-00041.safetensors", + "size_bytes": 3999853216, + "sha256": "a7794475040ebd62c9a1f9c94c17e7c600873e14fcabc656b32faea09bc7fd2d" + }, + { + "filename": "model-00007-of-00041.safetensors", + "size_bytes": 3999841912, + "sha256": "64d7e90d00ce15cc8bbf7677ad07bf3cffc9e90a1b777e3334fa15ffea219d6b" + }, + { + "filename": "model-00008-of-00041.safetensors", + "size_bytes": 3999842000, + "sha256": "9ace9b99e71490d619656956d457744f727647b3f36fa3d801798a5156599d35" + }, + { + "filename": "model-00009-of-00041.safetensors", + "size_bytes": 3999843192, + "sha256": "5b1b374c9e6100077c446a293566c1f644475c94739ab08b7fa26ba847110216" + }, + { + "filename": "model-00010-of-00041.safetensors", + "size_bytes": 3999517808, + "sha256": "38e51dc850a39c325324e6ddd23347e6e3e76cba1341aeef2e1f260ca2cb9f49" + }, + { + "filename": "model-00011-of-00041.safetensors", + "size_bytes": 4000181296, + "sha256": "c06719dc79fbbc796b8751ccbb29c8ecbdb686de9f89e3c741b5bbfec203cf0c" + }, + { + "filename": "model-00012-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "349036c3e3f6a8645af47906dae7272bd5352f822a3c354a16ef4084b3555679" + }, + { + "filename": "model-00013-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "bfab17d1535ef23e99e54ff3156ce1456aa5f0a4a93b281a7adb9d8b921dc829" + }, + { + "filename": "model-00014-of-00041.safetensors", + "size_bytes": 3999843984, + "sha256": "318d10df78f1647189941884acc01616a6a78d5d45d74ea30955be7f509c580a" + }, + { + "filename": "model-00015-of-00041.safetensors", + "size_bytes": 4000181736, + "sha256": "aa3400abde789ecca625b8cea37aa31232bf6696923f10981d028518a2393173" + }, + { + "filename": "model-00016-of-00041.safetensors", + "size_bytes": 3999517256, + "sha256": "3776700f68222f0174bb14fea8480da41d928a1bfd6b5317c4fc2b14006d506e" + }, + { + "filename": "model-00017-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "052726cd55224d3d217dac93eaed060215b8029afb44861c0d927d87c6a046ab" + }, + { + "filename": "model-00018-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "dd1b8d843e0013ed3beb6ba7fc36dc4b1f2df5d71b40520f82f369cd8fce5cba" + }, + { + "filename": "model-00019-of-00041.safetensors", + "size_bytes": 3999844096, + "sha256": "bfe5489c61334f9ce9bc6b369c3edf44420d7f585c6438f3afb98a120b26e6f3" + }, + { + "filename": "model-00020-of-00041.safetensors", + "size_bytes": 3999855040, + "sha256": "eba3113cbd34304751e148502f314be5ceae6e7d830bd33f2c0cf3bb8dfb28c2" + }, + { + "filename": "model-00021-of-00041.safetensors", + "size_bytes": 3999843792, + "sha256": "9ddd72e8ada4480e86e4230639210df1039802db9f78537673c19117d2e69b95" + }, + { + "filename": "model-00022-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "c6fc2bb1762d34b8db83fb11566e462c9166e132bb2bfcef45a33a4bd41c02db" + }, + { + "filename": "model-00023-of-00041.safetensors", + "size_bytes": 3999517464, + "sha256": "484bf8c327f7f7aaa20fbc4f6800d4aab7870749335b6f6bbfdb83570a1f66e2" + }, + { + "filename": "model-00024-of-00041.safetensors", + "size_bytes": 3999844264, + "sha256": "cfbb94709f5dac71ffe144ee3b9a976a49a2a9b9863caa29620fbc65820c2342" + }, + { + "filename": "model-00025-of-00041.safetensors", + "size_bytes": 4000181296, + "sha256": "b5799d19dcccfb17b108b5349900a4c42eb7bc241149cb7d1b38b48c707070cf" + }, + { + "filename": "model-00026-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "98003819929196d2f6f05c42603e9124a7a7853f372089eed3ffbd0c244fbdac" + }, + { + "filename": "model-00027-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "3d6b8b9d1c0c262cc8a63c41c3f75053ab8b31b6107130800718048da64b86db" + }, + { + "filename": "model-00028-of-00041.safetensors", + "size_bytes": 3999843984, + "sha256": "836d4a85cc72243c6caee2b6cf470cd598d6ab817e0685daf1128fd97bb133e3" + }, + { + "filename": "model-00029-of-00041.safetensors", + "size_bytes": 3999855320, + "sha256": "7b083d02c732458cfaaee8b0d1407626ad5ee10fc7eeed6b294b7d6853ac30e8" + }, + { + "filename": "model-00030-of-00041.safetensors", + "size_bytes": 3999843672, + "sha256": "e81006eeaa1152c67798640be46b710e87d44a307effca60cbdb7e9d1a3b268b" + }, + { + "filename": "model-00031-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "4fd46dbc55ceefa6146a7687cd02d040087c58b0f7ea42c4cd7dc66b838f531c" + }, + { + "filename": "model-00032-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "7edf6738b9bc204f0b39e9c7899515d1a1ed9d7d5f7c1cce32c981120c4f034d" + }, + { + "filename": "model-00033-of-00041.safetensors", + "size_bytes": 3999517688, + "sha256": "9f0f735e75e8d648b4b8d601a598d489fc4e4c7876da4f3406bf3bcc8990a181" + }, + { + "filename": "model-00034-of-00041.safetensors", + "size_bytes": 4000181496, + "sha256": "3499ba70dbb6a9acf87509c3b82c5790461a73f8d0d12f0e3699354f34dc02b7" + }, + { + "filename": "model-00035-of-00041.safetensors", + "size_bytes": 3999843792, + "sha256": "de04b5e39b9494b44becacf22fbe9c02890ef323a3dfe37640a2bfa80335504c" + }, + { + "filename": "model-00036-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "3eddb88c62194513f94ef3b1c11577ad324fdcc21265ff7b11a4b50d50b03c5f" + }, + { + "filename": "model-00037-of-00041.safetensors", + "size_bytes": 3999843872, + "sha256": "01992493e00d45d5fec4b5989392f401c9ebe3e13973edaaf1973577e132ca11" + }, + { + "filename": "model-00038-of-00041.safetensors", + "size_bytes": 3999844264, + "sha256": "e328c6d2e5f2a42eb99a0335bde7aa4c42cbd72b88fb988d5a2c96ba6231e95e" + }, + { + "filename": "model-00039-of-00041.safetensors", + "size_bytes": 3999854888, + "sha256": "15063cd0a46af3fa37e74590534f62a81f1ede5340e7640713bf98f39f20961d" + }, + { + "filename": "model-00040-of-00041.safetensors", + "size_bytes": 3365572496, + "sha256": "e458c4f4345d7082e7063f381ee1962725f84c44d53a5b70afb77e601ff22017" + }, + { + "filename": "model-00041-of-00041.safetensors", + "size_bytes": 3301131296, + "sha256": "1f3ac4d828f7e08dd14eb4dc6f282139ae1fa0d894e43f247da2539d2ef43826" + } + ], + "weight_files_total_size_bytes": 162659161528 + } + }, + "chain_semantics": { + "execution_dtype": "bfloat16", + "reference_dtype": "float32", + "accumulation_dtype": "float32", + "temperature": 1.0, + "loss_reduction": "sum_over_active_tokens_then_optional_mean_by_active_count", + "logprob_selection": "selected_token_logprob_on_active_mask", + "active_token_policy": "active selected tokens only", + "aggregates": [ + "max_abs_dlogp", + "approx_kl0", + "clipfrac0" + ], + "clip_interval": [ + 0.8, + 1.2 + ], + "clip_interval_note": "Pinned for clipfrac0; must match C1 chain_logprob_aggregates.default_clip_interval unless an explicit contract revision changes both.", + "comparison_roles_source": "rl_engine/kernels/gtest/tolerance_contract.json", + "forbidden_comparison_roles": [ + "baseline", + "singleton_aggregate" + ], + "singleton_aggregate_note": "singleton_aggregate is a C2 execution/aggregation mode only. It must never populate comparison_lhs_role or comparison_rhs_role.", + "tf32_policy_ref": "rl_engine/kernels/gtest/tolerance_contract.json#/policy/tf32", + "tf32_note": "WS1 TF32 enable/disable is owned by the C1 contract; C2 gates must not introduce a private TF32 policy.", + "report_naming": { + "comparison_lhs_role": "from_c1_by_report_kind", + "comparison_rhs_role": "from_c1_by_report_kind", + "forbidden_in_reports": [ + "baseline", + "singleton_aggregate" + ], + "singleton_aggregate_is": "c2_execution_aggregation_mode_only", + "note": "C2 freezes naming rules; C3+ emit reports that must obey these roles." + }, + "backend_actual_semantics": { + "c2_representative_actual_source": "scripts/ws1_candidate_evidence.py runtime execution", + "full_model_runtime_observed_actual_owner": [ + "C3", + "C8", + "C10", + "C11" + ], + "note": "C2 executes every representative case and records runtime-observed actual backend/kernel provenance. Later children own full-model dispatch provenance." + } + }, + "stochastic_policy": { + "dropout": 0.0, + "attention_dropout": 0.0, + "sampling_in_logprob_parity": false, + "canonical_gate_uses_dropout_zero": true, + "rng_source": "manifest_seed_plus_logical_sample_token_identity", + "undeclared_randomness": "hard_fail", + "retained_stochastic_ops": [] + }, + "primary_matrix": { + "description": "Fixed #150 Batch \u00d7 Chunked-Prefill matrix prerequisite workload cells.", + "N": 4, + "batch_size_bn": 4, + "sample_ids": [ + "s0", + "s1", + "s2", + "s3" + ], + "sample_order_fixed": true, + "batch_permutation": { + "enabled": true, + "permutation": [ + 2, + 0, + 3, + 1 + ], + "target_sample_position_in_bn": 0, + "note": "Permutation exercises layout invariance; logical compare restores sample_id order." + }, + "chunk": { + "chunk_size_tokens": 7, + "require_ge_2_chunks": true, + "non_divisible_case": true, + "note": "Longest primary seq_len=19 with chunk_size=7 yields chunks [7,7,5]." + }, + "cells": [ + { + "cell_id": "B1-singleton_aggregate/full", + "batch_mode": "singleton_aggregate", + "batch_size_per_run": 1, + "num_runs": 4, + "prefill_mode": "full", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "BN/full", + "batch_mode": "batched", + "batch_size_per_run": 4, + "num_runs": 1, + "prefill_mode": "full", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "B1-singleton_aggregate/chunked", + "batch_mode": "singleton_aggregate", + "batch_size_per_run": 1, + "num_runs": 4, + "prefill_mode": "chunked", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "BN/chunked", + "batch_mode": "batched", + "batch_size_per_run": 4, + "num_runs": 1, + "prefill_mode": "chunked", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + } + ] + }, + "fixtures": { + "prompt_template": "ws1_fixed_token_fixture", + "dtype_for_token_tensors": "int64", + "position_ids": { + "basis": "logical_zero_based_per_sample", + "reset_after_pack_boundary": true + }, + "attention_mask": { + "active_value": 1, + "padding_value": 0, + "causal": true + }, + "primary_seq_len": 19, + "primary_prompt_len": 8, + "short_seq_len": 8, + "long_seq_len": 32, + "varlen_seq_lens": [ + 11, + 16, + 13, + 19 + ], + "padding": { + "modes": [ + "right", + "left" + ], + "pad_token_id": 151643, + "primary_padded_len": 20 + }, + "packing": { + "status": "supported", + "implementation": "rl_engine.kernels.ops.pytorch.packing.pack.NativePackOp", + "packed_fixture": { + "sample_order": [ + "s0", + "s1", + "s2", + "s3" + ], + "segment_lengths": [ + 11, + 16, + 13, + 19 + ], + "total_tokens": 59, + "restore_key": [ + "sample_id", + "token_position" + ] + } + }, + "loss_mask": { + "prompt_tokens_active": false, + "completion_tokens_active": true + }, + "samples": [ + { + "sample_id": "s0", + "seq_len": 11, + "prompt_len": 8, + "token_ids": [ + 100, + 101, + 102, + 103, + 104, + 105, + 106, + 107, + 200, + 201, + 202 + ] + }, + { + "sample_id": "s1", + "seq_len": 16, + "prompt_len": 8, + "token_ids": [ + 110, + 111, + 112, + 113, + 114, + 115, + 116, + 117, + 210, + 211, + 212, + 213, + 214, + 215, + 216, + 217 + ] + }, + { + "sample_id": "s2", + "seq_len": 13, + "prompt_len": 8, + "token_ids": [ + 120, + 121, + 122, + 123, + 124, + 125, + 126, + 127, + 220, + 221, + 222, + 223, + 224 + ] + }, + { + "sample_id": "s3", + "seq_len": 19, + "prompt_len": 8, + "token_ids": [ + 130, + 131, + 132, + 133, + 134, + 135, + 136, + 137, + 230, + 231, + 232, + 233, + 234, + 235, + 236, + 237, + 238, + 239, + 240 + ] + } + ], + "short_full_model_fixture": { + "fixture_id": "short_full_model_seq8", + "seq_len": 8, + "prompt_len": 4, + "token_ids": [ + 310, + 311, + 312, + 313, + 410, + 411, + 412, + 413 + ], + "note": "Shorter sequence on full architecture+weights only; never shrinks layers/hidden/heads/vocab.", + "candidate_case_ids": [ + "short_full_model_seq8_qwen3_next_rms_norm", + "short_full_model_seq8_rms_norm_gated" + ] + }, + "long_full_model_fixture": { + "fixture_id": "long_full_model_seq32", + "seq_len": 32, + "prompt_len": 16, + "token_ids": [ + 500, + 501, + 502, + 503, + 504, + 505, + 506, + 507, + 508, + 509, + 510, + 511, + 512, + 513, + 514, + 515, + 600, + 601, + 602, + 603, + 604, + 605, + 606, + 607, + 608, + 609, + 610, + 611, + 612, + 613, + 614, + 615 + ], + "note": "Long fixed sequence on the same full architecture and pinned weight snapshot.", + "candidate_case_ids": [ + "long_full_model_seq32_qwen3_next_rms_norm", + "long_full_model_seq32_rms_norm_gated" + ] + }, + "representative_full_model_fixture": { + "fixture_id": "rep_full_model_seq16", + "seq_len": 16, + "prompt_len": 8, + "sample_ids": [ + "s0", + "s1", + "s2", + "s3" + ], + "note": "Primary variable-length matrix fixture; full architecture+weights.", + "candidate_case_ids": [ + "rep_full_model_seq16_qwen3_next_rms_norm", + "rep_full_model_seq16_rms_norm_gated" + ] + }, + "prompt_lens": [ + 8, + 8, + 8, + 8 + ], + "completion_lens": [ + 3, + 8, + 5, + 11 + ], + "max_completion_len": 11 + }, + "logical_identity": { + "key": [ + "sample_id", + "token_position" + ], + "token_position_basis": "logical_unpadded_index_in_sample", + "restore_before_compare_after": [ + "pad", + "pack", + "chunk", + "batch_permute" + ], + "gradient_singleton_aggregate": { + "definition": "N independent B=1 runs of the same N logical samples, aggregated with fixed sample order and active-token denominator", + "compare_to": "single B=N run of the same logical sample/token multiset", + "forbid_different_sample_sets": true + } + }, + "capabilities": { + "required_chain_ops": [ + { + "op": "qwen3_next_rms_norm", + "status": "required" + }, + { + "op": "rms_norm_gated", + "status": "required" + } + ], + "operator_spec_map": { + "qwen3_next_rms_norm": "qwen3_next_rms_norm", + "rms_norm_gated": "rms_norm_gated" + } + }, + "backend_profiles": { + "cuda_bf16": { + "backend_family": "cuda", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "declared", + "expected_backend_id": "cuda", + "expected_kernel_config_id": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp", + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "declared", + "expected_backend_id": "cuda", + "expected_kernel_config_id": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp", + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + }, + "triton_cuda_bf16": { + "backend_family": "triton", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + }, + "ascend_bf16": { + "backend_family": "ascend", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + } + }, + "representative_cases": [ + { + "case_id": "short_full_model_seq8_qwen3_next_rms_norm", + "fixture_id": "short_full_model_seq8", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "short_full_model_seq8_rms_norm_gated", + "fixture_id": "short_full_model_seq8", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "long_full_model_seq32_qwen3_next_rms_norm", + "fixture_id": "long_full_model_seq32", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "long_full_model_seq32_rms_norm_gated", + "fixture_id": "long_full_model_seq32", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "rep_full_model_seq16_qwen3_next_rms_norm", + "fixture_id": "rep_full_model_seq16", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "rep_full_model_seq16_rms_norm_gated", + "fixture_id": "rep_full_model_seq16", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + } + ], + "fixture_identity_sha256": "3ac9490ffbb1486edac6ea509811fd4d5affbe880b208bd7d9b052b0ce418442", + "provenance_boundary": { + "scope": "Synthetic norm inputs at checkpoint dimensions; not checkpoint execution or model-level L2.", + "runtime_verified": false + }, + "scope": "qwen3_next_norm_operators", + "full_model_evidence": false +} diff --git a/rl_engine/validation/models/qwen3_next_workload.py b/rl_engine/validation/models/qwen3_next_workload.py new file mode 100644 index 000000000..ee56f90da --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_workload.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Qwen3-Next norm-only C3/C4 workload, explicitly separate from the Dense chain.""" + +from collections.abc import Mapping +from typing import Any + +from rl_engine.config.workload import ( + WorkloadError, + _validate_backend_profiles, + _validate_fixtures, + _validate_logical_identity, + _validate_model_identity, + _validate_primary_matrix, + _validate_stochastic_policy, + manifest_identity_hash, +) + +MODEL_ID = "Qwen/Qwen3-Next-80B-A3B-Instruct" +REVISION = "9c7f2fbe84465e40164a94cc16cd30b6999b0cc7" +FINGERPRINT = { + "num_hidden_layers": 48, + "hidden_size": 2048, + "intermediate_size": 5120, + "num_attention_heads": 16, + "num_key_value_heads": 2, + "head_dim": 256, + "vocab_size": 151936, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "num_experts": 512, + "num_experts_per_tok": 10, +} +NORM_OPS = {"qwen3_next_rms_norm": 2048, "rms_norm_gated": 128} + + +def validate_norm_manifest(raw: Mapping[str, Any]) -> None: + identity = raw["model_identity"] + if identity["model_id"] != MODEL_ID or identity["revision"] != REVISION: + raise WorkloadError("Qwen3-Next norm workload requires the pinned official checkpoint") + _validate_model_identity(identity, fingerprint=FINGERPRINT, model_label="Qwen3-Next") + if ( + raw.get("scope") != "qwen3_next_norm_operators" + or raw.get("full_model_evidence") is not False + ): + raise WorkloadError("norm workload cannot claim full-model evidence") + _validate_stochastic_policy(raw["stochastic_policy"]) + _validate_primary_matrix(raw["primary_matrix"], raw["fixtures"]) + _validate_fixtures(raw["fixtures"], raw["primary_matrix"]) + _validate_logical_identity(raw["logical_identity"]) + caps = raw["capabilities"] + if {e["op"] for e in caps["required_chain_ops"]} != set(NORM_OPS): + raise WorkloadError("norm workload must contain exactly the two Qwen3-Next norms") + if any(e["status"] != "required" for e in caps["required_chain_ops"]): + raise WorkloadError("both norm operators are required") + _validate_backend_profiles(raw["backend_profiles"], caps) + cases = {case["case_id"]: case for case in raw["representative_cases"]} + if len(cases) != len(raw["representative_cases"]): + raise WorkloadError("duplicate norm case ID") + referenced = set() + for key in ( + "short_full_model_fixture", + "long_full_model_fixture", + "representative_full_model_fixture", + ): + fixture = raw["fixtures"][key] + for case_id in fixture["candidate_case_ids"]: + if case_id not in cases or cases[case_id]["fixture_id"] != fixture["fixture_id"]: + raise WorkloadError("norm case fixture binding mismatch") + referenced.add(case_id) + if referenced != set(cases): + raise WorkloadError("unreferenced norm case") + if {case["operator_spec"] for case in cases.values()} != set(NORM_OPS): + raise WorkloadError("representative cases must cover both norms") + for case in cases.values(): + if case["hidden"] != NORM_OPS[case["operator_spec"]]: + raise WorkloadError("norm case hidden dimension does not match checkpoint") + if case["architecture_identity"] != "qwen3_next_80b_a3b_norm_operators": + raise WorkloadError("norm cases cannot claim Dense or full-model architecture evidence") + if raw["fixture_identity_sha256"] != manifest_identity_hash(raw): + raise WorkloadError("Qwen3-Next fixture identity hash mismatch") + + +def validate_norm_dimensions(raw: Mapping[str, Any], op: str, hidden: int, head_dim: int) -> None: + if raw.get("scope") != "qwen3_next_norm_operators": + return + if op not in NORM_OPS or hidden != 2048 or head_dim != 128: + raise WorkloadError( + "Qwen3-Next norm gate requires its two norm ops, --hidden 2048 --head-dim 128" + ) diff --git a/rl_engine/validation/operators/gradient_adapters.py b/rl_engine/validation/operators/gradient_adapters.py index f22d80285..9831fb670 100644 --- a/rl_engine/validation/operators/gradient_adapters.py +++ b/rl_engine/validation/operators/gradient_adapters.py @@ -53,6 +53,7 @@ class GradientAdapterSpec: source_files: tuple[str, ...] shape_dependent_bwd_accum: str = "forbidden" atomic_add: str = "forbidden" + model_id: str | None = None @dataclass(frozen=True) @@ -115,6 +116,26 @@ def to_dict(self) -> dict[str, Any]: "csrc/cuda/norm/rmsnorm.cu", ), ), + "qwen3_next_rms_norm": GradientAdapterSpec( + op_name="qwen3_next_rms_norm", + chain_node="qwen3_next_rms_norm", + op_class="reduction", + spec_name="qwen3_next_rms_norm", + tensors=(_DX, _DWEIGHT), + requirement="required", + source_files=("rl_engine/backends/cuda/norm/rmsnorm.py", "csrc/cuda/norm/rmsnorm.cu"), + model_id="Qwen/Qwen3-Next-80B-A3B-Instruct", + ), + "rms_norm_gated": GradientAdapterSpec( + op_name="rms_norm_gated", + chain_node="rms_norm_gated", + op_class="reduction", + spec_name="rms_norm_gated", + tensors=(_DX, _DWEIGHT, _DGATE), + requirement="required", + source_files=("rl_engine/backends/cuda/norm/rmsnorm.py", "csrc/cuda/norm/rmsnorm.cu"), + model_id="Qwen/Qwen3-Next-80B-A3B-Instruct", + ), "qk_norm": GradientAdapterSpec( op_name="qk_norm", chain_node="qk_norm", @@ -291,18 +312,27 @@ def get_adapter(op_name: str) -> GradientAdapterSpec: raise KeyError(f"unknown gradient adapter {op_name!r}") from exc -def required_gradient_adapters() -> tuple[GradientAdapterSpec, ...]: +def required_gradient_adapters( + manifest: WS1Manifest | None = None, +) -> tuple[GradientAdapterSpec, ...]: + selected = manifest or load_manifest() + model_id = selected.model_identity["model_id"] + subset = selected.raw.get("scope") == "qwen3_next_norm_operators" return tuple( spec for spec in GRADIENT_ADAPTERS.values() if spec.requirement in ("required", "layout_supported") + and spec.model_id in (None, model_id) + and (not subset or spec.model_id == model_id) ) -def required_forward_adapters() -> tuple[GradientAdapterSpec, ...]: +def required_forward_adapters( + manifest: WS1Manifest | None = None, +) -> tuple[GradientAdapterSpec, ...]: """Same enumerable WS1 ops as C4; C3 reuses the registry, not a second list.""" - return required_gradient_adapters() + return required_gradient_adapters(manifest) @dataclass(frozen=True) @@ -620,9 +650,9 @@ def _row_parameters( head_dim: int = 16, ) -> dict[str, torch.Tensor]: """Config-independent trainable parameters, built in the execution dtype.""" - if op_name == "rms_norm": + if op_name in {"rms_norm", "qwen3_next_rms_norm"}: return {"weight": _shared_parameter((hidden,), device=device, dtype=dtype, offset=1)} - if op_name == "qk_norm": + if op_name in {"qk_norm", "rms_norm_gated"}: return {"weight": _shared_parameter((head_dim,), device=device, dtype=dtype, offset=1)} if op_name == "det_gemm": return {"b": _shared_parameter((hidden, hidden), device=device, dtype=dtype, offset=2)} @@ -663,12 +693,19 @@ def _row_inputs( """ n = len(keys) leading = (n,) - if op_name == "rms_norm": + if op_name in {"rms_norm", "qwen3_next_rms_norm"}: return { "x": _stack_rows(keys, leading, (hidden,), device=device, dtype=dtype), "weight": params["weight"], "eps": 1.0e-6, } + if op_name == "rms_norm_gated": + return { + "x": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype), + "gate": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype, offset=7), + "weight": params["weight"], + "eps": 1.0e-6, + } if op_name == "qk_norm": return { "x": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype), @@ -1228,6 +1265,15 @@ def resolve_profile_candidate( manifest: WS1Manifest | None = None, ) -> dict[str, Any]: m = manifest if manifest is not None else load_manifest() + subset = m.raw.get("scope") == "qwen3_next_norm_operators" + if adapter.model_id not in (None, m.model_identity["model_id"]) or ( + subset and adapter.model_id != m.model_identity["model_id"] + ): + return { + "status": "absent_not_required", + "expected_backend_id": None, + "candidate_path": None, + } if adapter.requirement == "absent_not_required": return { "status": "absent_not_required", diff --git a/rl_engine/validation/operators/operator_inputs.py b/rl_engine/validation/operators/operator_inputs.py index ca3b7c120..1cd03d31f 100644 --- a/rl_engine/validation/operators/operator_inputs.py +++ b/rl_engine/validation/operators/operator_inputs.py @@ -26,6 +26,8 @@ def make_operator_inputs( ) -> dict[str, Any]: builders = { "rms_norm": _make_rms_norm_inputs, + "rms_norm_gated": _make_rms_norm_gated_inputs, + "qwen3_next_rms_norm": _make_rms_norm_inputs, "qk_norm": _make_qk_norm_inputs, "pack": _make_pack_inputs, "matmul": _make_matmul_inputs, @@ -54,6 +56,8 @@ def operator_shape_name(op_name: str, args: argparse.Namespace) -> str: vocab = _arg_int(args, "vocab", DEFAULT_VOCAB) names = { "rms_norm": f"{batch}x{seq}x{_normalized_dim(args)}", + "rms_norm_gated": f"{batch}x{seq}x{_arg_int(args, 'head_dim', DEFAULT_HEAD_DIM)}", + "qwen3_next_rms_norm": f"{batch}x{seq}x{_normalized_dim(args)}", "qk_norm": f"{batch}x{seq}x{_arg_int(args, 'n_heads', DEFAULT_N_HEADS)}x" f"{_arg_int(args, 'head_dim', DEFAULT_HEAD_DIM)}", "pack": f"{batch}x{seq}x{_normalized_dim(args)}", @@ -91,6 +95,25 @@ def _make_rms_norm_inputs( } +def _make_rms_norm_gated_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + """Gated RMSNorm inputs at the GDN block's width. + + The gated norm normalizes over ``linear_value_head_dim`` (128 for + Qwen3-Next), not ``hidden_size``, so it follows ``head_dim`` rather than + ``normalized_dim``. + """ + batch, seq = _batch_seq(args) + head_dim = _arg_int(args, "head_dim", DEFAULT_HEAD_DIM) + return { + "x": _floating_tensor((batch, seq, head_dim), args, dtype, device, offset=0), + "weight": _floating_tensor((head_dim,), args, dtype, device, offset=1), + "gate": _floating_tensor((batch, seq, head_dim), args, dtype, device, offset=2), + "eps": _arg_float(args, "eps", DEFAULT_RMS_EPS), + } + + def _make_qk_norm_inputs( args: argparse.Namespace, dtype: torch.dtype, device: torch.device ) -> dict[str, Any]: diff --git a/rl_engine/validation/operators/operator_specs.py b/rl_engine/validation/operators/operator_specs.py index 2323f4787..7f17c369d 100644 --- a/rl_engine/validation/operators/operator_specs.py +++ b/rl_engine/validation/operators/operator_specs.py @@ -46,6 +46,30 @@ def _load_object(path: str) -> Any: }, grad_input_names=("x", "weight"), ), + "rms_norm_gated": OperatorSpec( + name="rms_norm_gated", + op_class="reduction", + gold_path=("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp"), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp"), + "cuda": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp"), + "cuda-sm90": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp"), + }, + grad_input_names=("x", "weight", "gate"), + ), + "qwen3_next_rms_norm": OperatorSpec( + name="qwen3_next_rms_norm", + op_class="reduction", + gold_path=("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp"), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp"), + "cuda": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp", + "cuda-sm90": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp"), + }, + grad_input_names=("x", "weight"), + ), "qk_norm": OperatorSpec( name="qk_norm", op_class="reduction", 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_gdn_recurrent_golden.py b/tests/models/qwen3_next/check_gdn_recurrent_golden.py new file mode 100644 index 000000000..043df421a --- /dev/null +++ b/tests/models/qwen3_next/check_gdn_recurrent_golden.py @@ -0,0 +1,521 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The GDN recurrent-step golden against the kernel vLLM actually runs (RFC #428 C6). + +Named ``check_`` rather than ``test_``, following ``tests/distributed/check_*.py``: +this module imports real vLLM. The integration isolation test +asserts ``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_gdn_recurrent_golden.py -v + +The provider here is ``fused_recurrent_gated_delta_rule_packed_decode``, not +``fused_sigmoid_gating_delta_rule_update``. A pure RL rollout decode -- N +sequences each emitting one token, no speculative decoding -- returns early into +the packed path because ``VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE`` defaults to +true. Aligning against the sigmoid-gating kernel would be validating a path +production does not take. + +Claim levels (RFC #428 section 2.1): L0 and L1 for the golden. L2 is not claimed; +the golden agrees with the provider at fp32-ULP scale but not bitwise, because +its contractions run in a fixed 32-wide chunk order rather than the kernel's +tree. + +The fixed-seed 128-step rounding test below is a bounded synthetic regression, +not evidence of a universal drift plateau or model-logit agreement. Earlier +1024-step and prefill tables had no reproducible runner and are withdrawn until +those experiments are checked in. Operator outputs are not model logits. + +""" + +from __future__ import annotations + +import pytest +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") + +from rl_engine.reference.linear_attn import GatedDeltaRuleRecurrentStepOp # noqa: E402 +from rl_engine.validation.models.qwen3_next_workload import FINGERPRINT # noqa: E402 + +# Qwen3-Next-80B-A3B-Instruct dims, from the pinned checkpoint fingerprint. +_H = FINGERPRINT["linear_num_key_heads"] +_HV = FINGERPRINT["linear_num_value_heads"] +_K = FINGERPRINT["linear_key_head_dim"] +_V = FINGERPRINT["linear_value_head_dim"] +_SCALE = _K**-0.5 +_PACKED_DIM = _H * _K * 2 + _HV * _V + +# Every bound in this file is a regression bound against the vLLM provider (or, for +# the bf16-state drift test, against the golden itself), set from measurements. None +# is gate evidence, and none goes through rl_engine.contracts.numerical's +# resolve_tolerance. For scale, the gate contract's +# forward_accuracy/by_op_class/reduction row in +# rl_engine/contracts/profiles/precision/ws1.json is atol = rtol = 1e-4 for float32 +# and atol = 5e-2, rtol = 2e-2 for bfloat16. +# +# (max|d out|, max|d state|) per recurrent-state dtype. Read by +# tools/validation/models/ws1_gdn_provider_agreement.py, so both use the same bounds. +_RECURRENT_BOUNDS = { + # fp32 state: agreement is at fp32-ULP scale. + torch.float32: (1e-3, 1e-5), + # bf16 state: the store rounds every token, so the state carries a bf16 ULP. + torch.bfloat16: (1e-3, 5e-3), +} + + +def _vllm_step(): + pytest.importorskip("vllm", reason="vLLM is required to compare against the provider") + from vllm.third_party.flash_linear_attention.ops import ( + fused_recurrent_gated_delta_rule_packed_decode, + ) + + return fused_recurrent_gated_delta_rule_packed_decode + + +def _inputs(batch, num_blocks, state_dtype, io_dtype, seed, indices=None): + g = torch.Generator(device="cuda").manual_seed(seed) + rand = lambda *shape, dtype: torch.randn( # noqa: E731 + *shape, device="cuda", dtype=dtype, generator=g + ) + if indices is None: + indices = torch.arange(1, batch + 1, device="cuda", dtype=torch.int32) + return { + "mixed_qkv": rand(batch, _PACKED_DIM, dtype=io_dtype), + "a": rand(batch, _HV, dtype=io_dtype), + "b": rand(batch, _HV, dtype=io_dtype), + "A_log": rand(_HV, dtype=torch.float32), + "dt_bias": rand(_HV, dtype=torch.float32), + "state": rand(num_blocks, _HV, _V, _K, dtype=state_dtype) * 0.1, + "indices": indices, + } + + +def _run_provider(inp, io_dtype=torch.bfloat16): + """Returns (out, mutated_state). The kernel updates the state in place.""" + step = _vllm_step() + state = inp["state"].clone() + out = torch.empty(inp["mixed_qkv"].shape[0], 1, _HV, _V, device="cuda", dtype=io_dtype) + step( + inp["mixed_qkv"], + inp["a"], + inp["b"], + inp["A_log"], + inp["dt_bias"], + _SCALE, + state, + out, + inp["indices"], + use_qk_l2norm_in_kernel=True, + ) + return out, state + + +def _run_golden(inp): + return GatedDeltaRuleRecurrentStepOp().forward( + inp["mixed_qkv"], + inp["a"], + inp["b"], + inp["A_log"], + inp["dt_bias"], + inp["state"].clone(), + inp["indices"], + scale=_SCALE, + num_k_heads=_H, + ) + + +# --------------------------------------------------------------------------- # +# 1. Against the provider +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("batch", [1, 4, 17, 64]) +@pytest.mark.parametrize("state_dtype", list(_RECURRENT_BOUNDS), ids=["fp32", "bf16"]) +def test_golden_matches_packed_decode_provider(batch, state_dtype): + out_atol, state_atol = _RECURRENT_BOUNDS[state_dtype] + inp = _inputs(batch, batch + 2, state_dtype, torch.bfloat16, seed=batch) + out_ref, state_ref = _run_provider(inp) + out_got, state_got = _run_golden(inp) + + assert (out_got.float() - out_ref.float()).abs().max().item() <= out_atol + assert (state_got.float() - state_ref.float()).abs().max().item() <= state_atol + + +def test_null_block_id_is_skipped_by_both(): + """Index 0 means "no state": zeros out, and the block is left alone.""" + indices = torch.tensor([1, 0, 2, 0], device="cuda", dtype=torch.int32) + inp = _inputs(4, 4, torch.float32, torch.bfloat16, seed=7, indices=indices) + + out_ref, state_ref = _run_provider(inp) + out_got, state_got = _run_golden(inp) + + for row in (1, 3): + assert bool((out_ref[row] == 0).all()), f"provider row {row}" + assert bool((out_got[row] == 0).all()), f"golden row {row}" + assert torch.equal(state_got[0], inp["state"][0]), "block 0 must be untouched" + assert torch.equal(state_ref[0], inp["state"][0]), "provider must leave block 0 untouched" + + +# --------------------------------------------------------------------------- # +# 2. L1 -- a sequence is unaffected by the others sharing the batch +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("state_dtype", [torch.float32, torch.bfloat16]) +def test_golden_is_batch_invariant(state_dtype): + """Bitwise: the same sequence, alone or in a batch of 64, at any position.""" + batch = 64 + inp = _inputs(batch, batch + 2, state_dtype, torch.bfloat16, seed=3) + full_out, full_state = _run_golden(inp) + + op = GatedDeltaRuleRecurrentStepOp() + for row in (0, 1, 31, 63): + alone = { + key: (value[row : row + 1] if key in ("mixed_qkv", "a", "b", "indices") else value) + for key, value in inp.items() + } + out, state = op.forward( + alone["mixed_qkv"], + alone["a"], + alone["b"], + alone["A_log"], + alone["dt_bias"], + alone["state"].clone(), + alone["indices"], + scale=_SCALE, + num_k_heads=_H, + ) + assert torch.equal(out[0], full_out[row]), f"out row {row}" + block = int(inp["indices"][row]) + assert torch.equal(state[block], full_state[block]), f"state block {block}" + + +# --------------------------------------------------------------------------- # +# 3. The state-dtype rounding is modelled, not skipped +# --------------------------------------------------------------------------- # +def test_bf16_state_rounding_stays_within_fixed_fixture_bound(): + """Bound this fixed 128-step fixture; no general contraction claim.""" + batch, steps = 8, 128 + inp = _inputs(batch, batch + 2, torch.float32, torch.bfloat16, seed=11) + op = GatedDeltaRuleRecurrentStepOp() + + # Both runs start from the SAME bf16-representable state, so the initial + # cast is a no-op and drift[0] can only come from a per-token store. + seed_state = inp["state"].to(torch.bfloat16).float() + state_fp32 = seed_state.clone() + state_bf16 = seed_state.to(torch.bfloat16).clone() + assert torch.equal(state_fp32, state_bf16.float()), "the two runs must start equal" + drift = [] + for step in range(steps): + g = torch.Generator(device="cuda").manual_seed(100 + step) + token = torch.randn(batch, _PACKED_DIM, device="cuda", dtype=torch.bfloat16, generator=g) + common = (inp["a"], inp["b"], inp["A_log"], inp["dt_bias"]) + _, state_fp32 = op.forward( + token, *common, state_fp32, inp["indices"], scale=_SCALE, num_k_heads=_H + ) + _, state_bf16 = op.forward( + token, *common, state_bf16, inp["indices"], scale=_SCALE, num_k_heads=_H + ) + scale = max(state_fp32.float().abs().max().item(), 1e-9) + drift.append((state_fp32.float() - state_bf16.float()).abs().max().item() / scale) + + assert drift[0] > 0.0, "a bf16 state must round on the very first store" + # Preserve the original regression bound for this fixed fixture. + # Regression bound for this fixture, golden against golden; not gate evidence. + assert max(drift) < 0.05, f"relative state drift reached {max(drift):.3e}" + # The back half must not be materially worse than the front half. + assert max(drift[steps // 2 :]) < 2.0 * max(drift[: steps // 2]) + 1e-3 + + +# --------------------------------------------------------------------------- # +# 4. The causal-conv1d state update, the other half of a decode step +# --------------------------------------------------------------------------- # +# The conv runs over the packed q|k|v channels; Qwen3-Next's linear_conv_kernel_dim is 4. +_CONV_DIM, _CONV_WIDTH = _PACKED_DIM, 4 + + +def _conv_update(): + pytest.importorskip("vllm", reason="vLLM is required to compare against the provider") + from vllm.model_executor.layers.mamba.ops.causal_conv1d import causal_conv1d_update + + return causal_conv1d_update + + +def _conv_inputs(batch, state_dtype, seed, with_bias=True, indices=None): + g = torch.Generator(device="cuda").manual_seed(seed) + rand = lambda *shape, dtype: torch.randn( # noqa: E731 + *shape, device="cuda", dtype=dtype, generator=g + ) + if indices is None: + indices = torch.arange(1, batch + 1, device="cuda", dtype=torch.int32) + return { + "x": rand(batch, _CONV_DIM, dtype=torch.bfloat16), + # The cache must have room for every index used, plus the null block. + "state": rand(batch + 2, _CONV_DIM, _CONV_WIDTH - 1, dtype=state_dtype), + "weight": rand(_CONV_DIM, _CONV_WIDTH, dtype=torch.bfloat16), + "bias": rand(_CONV_DIM, dtype=torch.bfloat16) if with_bias else None, + "indices": indices, + } + + +def _run_conv_pair(inp, activation="silu"): + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + state_ref, out_ref = inp["state"].clone(), torch.empty_like(inp["x"]) + _conv_update()( + inp["x"], + state_ref, + inp["weight"], + inp["bias"], + activation, + conv_state_indices=inp["indices"], + out=out_ref, + ) + out_got, state_got = CausalConv1dUpdateOp().forward( + inp["x"], + inp["state"].clone(), + inp["weight"], + inp["indices"], + bias=inp["bias"], + activation=activation, + ) + return (out_ref, state_ref), (out_got, state_got) + + +@pytest.mark.parametrize("batch", [1, 4, 17, 64]) +def test_conv_state_update_is_bitwise_exact(batch): + """The rolled window is what the next token consumes, so it must be exact. + + Holds for an fp32 and a bf16 cache alike -- the rolling is a copy, not a + computation, which is also why bias and activation are not varied here: they + cannot reach the state. + """ + for state_dtype in (torch.float32, torch.bfloat16): + inp = _conv_inputs(batch, state_dtype, seed=batch) + (_, state_ref), (_, state_got) = _run_conv_pair(inp) + assert torch.equal(state_got, state_ref), f"{state_dtype} batch={batch}" + + +@pytest.mark.parametrize("state_dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("activation", ["silu", None]) +@pytest.mark.parametrize("with_bias", [True, False]) +def test_conv_golden_is_batch_invariant(state_dtype, activation, with_bias): + """L1 for the conv golden, bitwise: a sequence's output row and state block + do not depend on which other sequences share the step. + + Every sequence of a batch of 257 is run alone; sub-batches of 2, 7 and 64 at + other offsets and a shuffled batch order are run too, with contiguous and + shuffled cache blocks, over three seeds. + """ + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + op = CausalConv1dUpdateOp() + batch = 257 + for seed in (3, 4, 5): + g = torch.Generator(device="cpu").manual_seed(seed) + perm = torch.randperm(batch, generator=g).cuda() + shuffled_blocks = (torch.randperm(batch, generator=g) + 1).to("cuda", torch.int32) + for indices in (None, shuffled_blocks): + inp = _conv_inputs(batch, state_dtype, seed, with_bias=with_bias, indices=indices) + + def run(rows, inp=inp): + return op.forward( + inp["x"][rows], + inp["state"].clone(), + inp["weight"], + inp["indices"][rows], + bias=inp["bias"], + activation=activation, + ) + + full_out, full_state = run(torch.arange(batch, device="cuda")) + picks = [torch.tensor([r], device="cuda") for r in range(batch)] + picks += [ + torch.arange(o, o + s, device="cuda") for s, o in ((2, 0), (7, 100), (64, 193)) + ] + picks.append(perm) + for rows in picks: + out, state = run(rows) + blocks = inp["indices"][rows].long() + where = f"seed={seed} rows[{int(rows[0])}..] n={rows.numel()}" + assert torch.equal(out, full_out[rows]), f"out {where}" + assert torch.equal(state[blocks], full_state[blocks]), f"state {where}" + + +@pytest.mark.parametrize("batch", [1, 4, 17, 64]) +def test_conv_output_matches_provider_with_fp32_cache(batch): + """An fp32 cache reproduces the provider except on a few output elements. + + Measured by tools/validation/models/ws1_gdn_provider_agreement.py (B200, acf38b6): 0 of 8192 + at B=1, 0 at B=4, 1 of 139264 at B=17, 5 of 524288 at B=64; max|diff| 3.9e-3. + The bound is on the magnitude and on a handful of elements rather than on a + tight rate -- + at small batches a single straddling element is already 1.2e-4 of the + tensor, which says nothing about accuracy. + """ + inp = _conv_inputs(batch, torch.float32, seed=batch) + (out_ref, _), (out_got, _) = _run_conv_pair(inp) + mismatch = int((out_got.float().view(torch.int32) != out_ref.float().view(torch.int32)).sum()) + # Regression bounds against the provider, not gate evidence (see the module note). + assert mismatch <= 32, f"{mismatch} of {out_got.numel()} elements differ" + assert (out_got.float() - out_ref.float()).abs().max().item() <= 1e-2 + + +@pytest.mark.parametrize("batch", [1, 17]) +def test_conv_output_bf16_cache_matches_rounded_product_path(batch): + """Product rounding is reproduced; activation ULP residuals remain allowed.""" + inp = _conv_inputs(batch, torch.bfloat16, seed=batch) + (out_ref, _), (out_got, _) = _run_conv_pair(inp) + # Regression bounds against the provider, not gate evidence (see the module note). + assert (out_got.float() - out_ref.float()).abs().max().item() <= 7e-2 + assert int((out_got != out_ref).sum()) <= 32 + + +def test_conv_null_block_id_is_skipped(): + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + indices = torch.tensor([1, 0, 2, 0], device="cuda", dtype=torch.int32) + inp = _conv_inputs(4, torch.float32, seed=5, with_bias=False, indices=indices) + out_got, state_got = CausalConv1dUpdateOp().forward( + inp["x"], inp["state"].clone(), inp["weight"], indices, bias=None + ) + for row in (1, 3): + assert bool((out_got[row] == 0).all()) + assert torch.equal(state_got[0], inp["state"][0]) + + +def test_conv_rejects_unknown_activation_and_bad_shapes(): + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + op = CausalConv1dUpdateOp() + inp = _conv_inputs(2, torch.float32, seed=1) + with pytest.raises(ValueError, match="activation must be"): + op.forward(inp["x"], inp["state"], inp["weight"], inp["indices"], activation="relu") + with pytest.raises(ValueError, match="conv_state tail must be"): + op.forward( + inp["x"], + inp["state"][..., :1], + inp["weight"], + inp["indices"], + activation=None, + ) + + +# --------------------------------------------------------------------------- # +# 5. forward_fp32, and the branches the provider comparison never reaches +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("state_dtype", [torch.float32, torch.bfloat16]) +def test_recurrent_forward_fp32_widens_the_state(state_dtype): + """forward_fp32 must run the recurrence without the per-token rounding. + + With a bf16 cache that means widening it -- the whole point of having the + method is to be able to compare against the rounded path. + """ + op = GatedDeltaRuleRecurrentStepOp() + inp = _inputs(4, 6, state_dtype, torch.bfloat16, seed=1) + out, state = op.forward_fp32( + inp["mixed_qkv"], + inp["a"], + inp["b"], + inp["A_log"], + inp["dt_bias"], + inp["state"].clone(), + inp["indices"], + scale=_SCALE, + num_k_heads=_H, + ) + assert out.dtype is torch.float32 and state.dtype is torch.float32 + + rounded, _ = _run_golden(inp) + if state_dtype is torch.bfloat16: + # The rounded path went through bf16; the fp32 path did not. + assert not torch.equal(rounded.float(), out) + + +def test_recurrent_call_matches_forward(): + op = GatedDeltaRuleRecurrentStepOp() + inp = _inputs(4, 6, torch.float32, torch.bfloat16, seed=2) + args = ( + inp["mixed_qkv"], + inp["a"], + inp["b"], + inp["A_log"], + inp["dt_bias"], + inp["state"].clone(), + inp["indices"], + ) + kwargs = dict(scale=_SCALE, num_k_heads=_H) + a_out, a_state = op(*args, **kwargs) + b_out, b_state = op.forward(*args, **kwargs) + assert torch.equal(a_out, b_out) and torch.equal(a_state, b_state) + + +def test_recurrent_without_in_kernel_l2norm_differs(): + """`use_qk_l2norm=False` is the prefill convention and must be reachable.""" + op = GatedDeltaRuleRecurrentStepOp() + inp = _inputs(4, 6, torch.float32, torch.bfloat16, seed=3) + common = ( + inp["mixed_qkv"], + inp["a"], + inp["b"], + inp["A_log"], + inp["dt_bias"], + ) + normed, _ = op.forward( + *common, inp["state"].clone(), inp["indices"], scale=_SCALE, num_k_heads=_H + ) + raw, _ = op.forward( + *common, + inp["state"].clone(), + inp["indices"], + scale=_SCALE, + num_k_heads=_H, + use_qk_l2norm=False, + ) + assert not torch.equal(normed, raw) + + +def test_conv_dim_first_false_matches_the_transposed_layout(): + """`dim_first=False` is vLLM's "SD" cache layout, not a dead branch.""" + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + op = CausalConv1dUpdateOp() + inp = _conv_inputs(4, torch.float32, seed=9) + ds_out, ds_state = op.forward( + inp["x"], inp["state"].clone(), inp["weight"], inp["indices"], bias=inp["bias"] + ) + sd_out, sd_state = op.forward( + inp["x"], + inp["state"].transpose(-1, -2).contiguous(), + inp["weight"], + inp["indices"], + bias=inp["bias"], + dim_first=False, + ) + assert torch.equal(ds_out, sd_out) + assert torch.equal(ds_state, sd_state.transpose(-1, -2)) + + +def test_conv_forward_fp32_and_call_entry_points(): + from rl_engine.reference.linear_attn import CausalConv1dUpdateOp + + op = CausalConv1dUpdateOp() + inp = _conv_inputs(4, torch.float32, seed=10) + args = (inp["x"], inp["state"].clone(), inp["weight"], inp["indices"]) + out, _ = op.forward(*args, bias=inp["bias"]) + called, _ = op(*args, bias=inp["bias"]) + fp32, _ = op.forward_fp32(*args, bias=inp["bias"]) + assert torch.equal(out, called) + assert fp32.dtype is torch.float32 + torch.testing.assert_close(fp32, out.float(), atol=1e-2, rtol=1e-2) + + +def test_conv_provider_preserves_bf16_product_cancellation(): + inp = _conv_inputs(1, torch.bfloat16, seed=1, with_bias=False) + inp["state"].zero_() + inp["state"][1, :, -1] = 1.0078125 + inp["weight"].zero_() + inp["weight"][:, -2] = 1.0078125 + inp["weight"][:, -1] = 1.0 + inp["x"].fill_(-1.015625) + (provider, _), (golden, _) = _run_conv_pair(inp, activation=None) + assert torch.count_nonzero(provider) == 0 + assert torch.equal(provider, golden) 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_gdn_state_contract.py b/tests/models/qwen3_next/test_gdn_state_contract.py new file mode 100644 index 000000000..c8aaec572 --- /dev/null +++ b/tests/models/qwen3_next/test_gdn_state_contract.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Provider-independent cache contracts, collected in the ordinary CPU suite.""" + +import pytest +import torch + +from rl_engine.reference.linear_attn import CausalConv1dUpdateOp, GatedDeltaRuleRecurrentStepOp + + +def _call(kind, indices, heads=1): + batch = indices.numel() + if kind == "conv": + state = torch.randn(4, 2, 3) + before = state.clone() + result = CausalConv1dUpdateOp()(torch.randn(batch, 2), state, torch.randn(2, 4), indices) + else: + state = torch.randn(4, 2, 2, 32) + before = state.clone() + result = GatedDeltaRuleRecurrentStepOp()( + torch.randn(batch, 68), + torch.randn(batch, 2), + torch.randn(batch, 2), + torch.zeros(2), + torch.zeros(2), + state, + indices, + scale=32**-0.5, + num_k_heads=heads, + ) + assert torch.equal(state, before) + return result, before + + +@pytest.mark.parametrize("kind", ["conv", "gdn"]) +@pytest.mark.parametrize( + "indices,message", + [ + (torch.tensor([1.5, 2.0]), "int32 or int64"), + (torch.tensor([True, False]), "int32 or int64"), + (torch.tensor([1, 1]), "unique"), + (torch.tensor([1, 4]), "out of range"), + ], +) +def test_invalid_cache_indices_are_rejected(kind, indices, message): + with pytest.raises(ValueError, match=message): + _call(kind, indices) + + +@pytest.mark.parametrize("kind", ["conv", "gdn"]) +@pytest.mark.parametrize("indices", [torch.tensor([0, -1, 0]), torch.empty(0, dtype=torch.int64)]) +def test_inactive_and_empty_batches_preserve_cache(kind, indices): + (out, state), before = _call(kind, indices) + assert torch.equal(state, before) + assert torch.count_nonzero(out) == 0 + + +@pytest.mark.parametrize("heads", [0, -1, 1.5, True]) +def test_gdn_requires_positive_integer_heads(heads): + with pytest.raises(ValueError, match="positive integer"): + _call("gdn", torch.tensor([1, 2]), heads=heads) + + +def test_conv_bias_is_accumulated_before_taps(): + # Adding bias last would yield 1.0 instead of 0.0. + out, _ = CausalConv1dUpdateOp()( + torch.tensor([[-1e8]]), + torch.zeros(2, 1, 1), + torch.tensor([[0.0, 1.0]]), + torch.tensor([1]), + bias=torch.tensor([1e8]), + activation=None, + ) + assert out.item() == 0.0 + out, _ = CausalConv1dUpdateOp()( + torch.tensor([[-1e8]]), + torch.ones(2, 1, 1), + torch.ones(1, 2), + torch.tensor([1]), + bias=torch.tensor([1e8]), + activation=None, + ) + assert out.item() == 0.0 + + +def test_conv_bf16_products_round_before_fp32_accumulation(): + state = torch.zeros(2, 1, 1, dtype=torch.bfloat16) + state[1, 0, 0] = 1.0078125 + out, _ = CausalConv1dUpdateOp().forward_fp32( + torch.tensor([[-1.015625]], dtype=torch.bfloat16), + state, + torch.tensor([[1.0078125, 1.0]], dtype=torch.bfloat16), + torch.tensor([1]), + activation=None, + ) + # The first exact product is 1.015686..., rounded to 1.015625 in BF16. + assert out.item() == 0.0 + + +def test_softplus_large_branch_has_finite_gradient(): + from rl_engine.reference.linear_attn.gated_delta_rule import _softplus + + x = torch.tensor([100.0, 21.0, 0.0], requires_grad=True) + _softplus(x).sum().backward() + assert torch.equal(x.grad, torch.tensor([1.0, 1.0, 0.5])) + + +def test_chunked_sum_keeps_fixed_order_for_non_multiple_of_32_width(): + from rl_engine.reference.linear_attn.gated_delta_rule import _REDUCTION_CHUNK, _chunked_sum + + torch.manual_seed(0) + x = torch.randn(3, 5, 2 * _REDUCTION_CHUNK + 7) + padded = torch.nn.functional.pad(x, (0, _REDUCTION_CHUNK - 7)) + chunked = padded.reshape(3, 5, 3, _REDUCTION_CHUNK).sum(dim=-1).sum(dim=-1) + assert torch.equal(_chunked_sum(x), chunked) + # The reduction must not depend on the leading shape either. + assert torch.equal(_chunked_sum(x[:1]), _chunked_sum(x)[:1]) 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..4497ff87f --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_norm.py @@ -0,0 +1,884 @@ +# 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 + + +# --------------------------------------------------------------------------- # +# 7. Gated RMSNorm CUDA kernel (the GDN block's norm) +# --------------------------------------------------------------------------- # +_HAS_CUDA_GATED = False +if torch.cuda.is_available(): # pragma: no branch - probe only + try: + from rl_engine.backends.extension import _C as _C_probe + from rl_engine.backends.extension import _EXT_AVAILABLE as _EXT_probe + + _HAS_CUDA_GATED = _EXT_probe and hasattr(_C_probe, "rmsnorm_gated_forward") + except ImportError: # pragma: no cover + _HAS_CUDA_GATED = False + +requires_cuda_gated = pytest.mark.skipif( + not _HAS_CUDA_GATED, reason="gated RMSNorm CUDA extension is not available" +) + + +def _gated_cuda_inputs(seed=0, rows=512, hidden=_HEAD_V_DIM, dtype=torch.bfloat16): + g = torch.Generator(device="cuda").manual_seed(seed) + x = torch.randn(rows, hidden, device="cuda", dtype=dtype, generator=g) + gate = torch.randn(rows, hidden, device="cuda", dtype=dtype, generator=g) + weight = torch.randn(hidden, device="cuda", dtype=dtype, generator=g) + return x, weight, gate + + +@requires_cuda_gated +def test_cuda_gated_matches_golden_within_contract(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + x, w, gate = _gated_cuda_inputs() + got = Qwen3NextRMSNormGatedCudaOp().forward(x, w, gate, eps=_EPS) + ref = Qwen3NextRMSNormGatedOp().forward_fp32(x, w, gate, eps=_EPS) + torch.testing.assert_close(got.float(), ref, **_forward_tol(torch.bfloat16)) + + +@requires_cuda_gated +@pytest.mark.parametrize("offset", [0.0, 1.0]) +@pytest.mark.parametrize("activation", [0, 1], ids=["silu", "sigmoid"]) +@pytest.mark.parametrize("hidden", [_HEAD_V_DIM, 2048, 5120]) +@pytest.mark.parametrize( + "dtype", [torch.float32, torch.float16, torch.bfloat16], ids=["fp32", "fp16", "bf16"] +) +def test_cuda_gated_rstd_is_bitwise_identical_to_plain_kernel(dtype, hidden, activation, offset): + """The gate must not perturb the normalization statistic. + + Same x, same rstd, bitwise -- otherwise the gate has leaked into the + reduction and the two kernels no longer share a contract. Both kernels + launch with ``choose_threads(H)``: H=128 runs 128 threads with one column + each, 2048 and 5120 run 512 threads with 4 and 10 serial columns each. + Each activation is its own template instantiation, and the offset is + applied only after the statistic, so neither it nor the gate may move it. + """ + from rl_engine.backends.extension import _C + + x, w, gate = _gated_cuda_inputs(hidden=hidden, dtype=dtype) + _, rstd_gated = _C.rmsnorm_gated_forward(x, w, gate, _EPS, offset, activation) + _, rstd_plain = _C.rmsnorm_forward(x, w, _EPS, offset) + assert torch.equal(rstd_gated, rstd_plain) + + +@requires_cuda_gated +@pytest.mark.parametrize("batch", [1, 2, 8, 16, 32, 48, 64, 512]) +def test_cuda_gated_batch_invariance(batch): + """L1 across the RFC #428 concurrency axis, bitwise.""" + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + op = Qwen3NextRMSNormGatedCudaOp() + x, w, gate = _gated_cuda_inputs() + full = op.forward(x, w, gate, eps=_EPS) + assert torch.equal(op.forward(x[:batch], w, gate[:batch], eps=_EPS), full[:batch]) + + +@requires_cuda_gated +def test_cuda_gated_zero_gate_zeroes_output(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + x, w, _ = _gated_cuda_inputs() + out = Qwen3NextRMSNormGatedCudaOp().forward(x, w, torch.zeros_like(x), eps=_EPS) + assert torch.equal(out, torch.zeros_like(out)) + + +@requires_cuda_gated +def test_cuda_gated_unit_weight_is_plain_norm_times_silu(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp, rmsnorm_cuda + + x, _, gate = _gated_cuda_inputs() + ones = torch.ones(_HEAD_V_DIM, device="cuda", dtype=x.dtype) + got = Qwen3NextRMSNormGatedCudaOp().forward(x, ones, gate, eps=_EPS).float() + ref = rmsnorm_cuda(x, ones, eps=_EPS).float() * F.silu(gate.float()) + torch.testing.assert_close(got, ref, **_forward_tol(torch.bfloat16)) + + +@requires_cuda_gated +def test_cuda_gated_sigmoid_activation(): + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_cuda, rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs() + got = rmsnorm_gated_cuda(x, w, gate, eps=_EPS, activation="sigmoid").float() + ref = rmsnorm_cuda(x, w, eps=_EPS).float() * torch.sigmoid(gate.float()) + torch.testing.assert_close(got, ref, **_forward_tol(torch.bfloat16)) + + +@requires_cuda_gated +def test_cuda_gated_backward_matches_autograd_golden(): + """dx, dweight and dgate against the fp32 reference's autograd.""" + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs() + x, w, gate = x.float(), w.float(), gate.float() + dy = torch.randn_like(x) + + got = [t.clone().requires_grad_(True) for t in (x, w, gate)] + rmsnorm_gated_cuda(*got, eps=_EPS).backward(dy) + ref = [t.clone().requires_grad_(True) for t in (x, w, gate)] + Qwen3NextRMSNormGatedOp().forward_fp32(*ref, eps=_EPS).backward(dy) + + for name, a, b in zip(("dx", "dweight", "dgate"), got, ref): + torch.testing.assert_close(a.grad, b.grad, atol=1e-4, rtol=1e-4, msg=name) + + +@requires_cuda_gated +@pytest.mark.parametrize("batch", [1, 8, 64, 512]) +def test_cuda_gated_dweight_is_batch_invariant(batch): + """A row's contribution to dweight must not depend on the batch it sits in. + + The ascending-row fp32 left fold makes dweight over the first n rows exactly + the fold of those n row-contributions -- so the n-row run must reproduce the + 512-row run's prefix, bitwise. + """ + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + from rl_engine.backends.extension import _C + from rl_engine.ops.autograd.vjp_fp32 import reduce_rows_fp32 + + x, w, gate = _gated_cuda_inputs() + x, w, gate = x.float(), w.float(), gate.float() + dy = torch.randn_like(x) + + wg = w.clone().requires_grad_(True) + rmsnorm_gated_cuda(x[:batch], wg, gate[:batch], eps=_EPS).backward(dy[:batch]) + + # rstd from the full-batch forward: the row statistic is batch-invariant, so + # its prefix is what the n-row run must have seen. Recomputing it with + # mean(-1) instead would compare against a different reduction. + _, rstd_full = _C.rmsnorm_gated_forward(x, w, gate, _EPS, 0.0, 0) + rows = dy * F.silu(gate) * x * rstd_full.unsqueeze(-1) + assert torch.equal(wg.grad, reduce_rows_fp32(rows[:batch])) + + +@requires_cuda_gated +def test_cuda_gated_weight_offset_is_applied_in_fp32(): + """The gated kernel carries the same offset contract as the plain one.""" + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs() + offset = rmsnorm_gated_cuda(x, w, gate, eps=_EPS, weight_offset=1.0) + explicit = rmsnorm_gated_cuda(x, (1.0 + w.float()), gate, eps=_EPS) + assert torch.equal(offset, explicit) + folded = rmsnorm_gated_cuda(x, (1.0 + w.float()).bfloat16(), gate, eps=_EPS) + assert not torch.equal(offset, folded) + + +@requires_cuda_gated +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("activation", ["silu", "sigmoid"]) +def test_cuda_gated_parameter_vjp_contributions_match_the_fold(dtype, activation): + """Compare the harness hook with actual autograd, including its CUDA rstd.""" + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp, _fold_dweight_rows + + x, w, gate = (t.to(dtype) for t in _gated_cuda_inputs(rows=512)) + w = w.float().requires_grad_() + dy = torch.randn_like(x) + op = Qwen3NextRMSNormGatedCudaOp(activation=activation) + op.forward(x, w, gate, eps=_EPS).backward(dy) + rows = op.parameter_vjp_contributions_fp32(x=x, weight=w, gate=gate, grad_output=dy, eps=_EPS)[ + "weight" + ] + assert torch.equal(_fold_dweight_rows(rows, torch.float32), w.grad) + + +@requires_cuda_gated +def test_cuda_gated_wrapper_rejects_equal_numel_wrong_shape(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + x, w, gate = _gated_cuda_inputs(rows=6) + x = x.reshape(2, 3, -1) + with pytest.raises(ValueError, match="gate must match x"): + Qwen3NextRMSNormGatedCudaOp().forward(x, w, gate) + + +@requires_cuda_gated +def test_cuda_gated_sigmoid_backward_uses_the_sigmoid_derivative(): + """The activation gradient has two branches; only silu was covered.""" + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs(rows=64) + x, w, gate = x.float(), w.float(), gate.float() + dy = torch.randn_like(x) + + got = [t.clone().requires_grad_(True) for t in (x, w, gate)] + rmsnorm_gated_cuda(*got, eps=_EPS, activation="sigmoid").backward(dy) + + ref = [t.clone().requires_grad_(True) for t in (x, w, gate)] + rstd = torch.rsqrt(ref[0].square().mean(dim=-1) + _EPS) + out = (ref[0] * rstd.unsqueeze(-1) * ref[1]) * torch.sigmoid(ref[2]) + out.backward(dy) + for name, a, b in zip(("dx", "dweight", "dgate"), got, ref): + torch.testing.assert_close(a.grad, b.grad, atol=1e-4, rtol=1e-4, msg=name) + + +@requires_cuda_gated +def test_cuda_gated_rejects_bad_activation(): + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs(rows=8) + with pytest.raises(ValueError, match="activation must be one of"): + rmsnorm_gated_cuda(x, w, gate, eps=_EPS, activation="gelu") + + +def test_gated_cuda_op_rejects_bad_activation_at_construction(): + """The activation is validated before the extension check, so this runs on CPU.""" + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + with pytest.raises(ValueError, match="activation must be one of"): + Qwen3NextRMSNormGatedCudaOp(activation="gelu") + + +@requires_cuda_gated +def test_cuda_gated_swish_is_an_alias_for_silu(): + """vLLM maps output_gate_type "swish" to "silu"; the alias must not select sigmoid.""" + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + x, w, gate = _gated_cuda_inputs(rows=64) + swish = Qwen3NextRMSNormGatedCudaOp(activation="swish").forward(x, w, gate, eps=_EPS) + silu = Qwen3NextRMSNormGatedCudaOp(activation="silu").forward(x, w, gate, eps=_EPS) + sigmoid = Qwen3NextRMSNormGatedCudaOp(activation="sigmoid").forward(x, w, gate, eps=_EPS) + assert torch.equal(swish, silu) + assert not torch.equal(swish, sigmoid) + + +@requires_cuda_gated +def test_cuda_gated_rejects_mismatched_gate(): + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x, w, gate = _gated_cuda_inputs(rows=8) + with pytest.raises((AssertionError, RuntimeError)): + rmsnorm_gated_cuda(x, w, gate[:4], eps=_EPS) + with pytest.raises((AssertionError, RuntimeError)): + rmsnorm_gated_cuda(x, w, gate.float(), eps=_EPS) + + +# --------------------------------------------------------------------------- # +# 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() + + +@requires_cuda_gated +@pytest.mark.parametrize( + "fault", ["weight_shape", "weight_dtype", "gate_dtype", "dy_dtype", "rstd_dtype"] +) +def test_gated_extension_rejects_invalid_tensor_contract(fault): + from rl_engine import _C + + x, w, gate = _gated_cuda_inputs(rows=8) + dy = torch.ones_like(x) + rstd = torch.ones(8, device=x.device, dtype=torch.float32) + if fault == "weight_shape": + w = w[:-1] + elif fault == "weight_dtype": + w = w.half() + elif fault == "gate_dtype": + gate = gate.float() + elif fault == "dy_dtype": + dy = dy.float() + else: + rstd = rstd.bfloat16() + with pytest.raises(RuntimeError): + _C.rmsnorm_gated_backward_dx(dy, x, w, gate, rstd) + + +@requires_cuda_gated +def test_gated_extension_handles_empty_batch_and_nondefault_stream(): + from rl_engine import _C + + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + x, w, gate = _gated_cuda_inputs(rows=8) + actual, rstd = _C.rmsnorm_gated_forward(x, w, gate, _EPS) + dx = _C.rmsnorm_gated_backward_dx(torch.ones_like(x), x, w, gate, rstd) + empty, empty_rstd = _C.rmsnorm_gated_forward(x[:0], w, gate[:0], _EPS) + empty_dx = _C.rmsnorm_gated_backward_dx(x[:0], x[:0], w, gate[:0], empty_rstd) + stream.synchronize() + expected, expected_rstd = _C.rmsnorm_gated_forward(x, w, gate, _EPS) + expected_dx = _C.rmsnorm_gated_backward_dx(torch.ones_like(x), x, w, gate, expected_rstd) + assert torch.equal(actual, expected) + assert torch.equal(dx, expected_dx) + assert empty.shape == empty_dx.shape == (0, x.shape[-1]) + assert empty_rstd.numel() == 0 + + +@requires_cuda_gated +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +def test_cuda_gated_signed_zero_weight_is_preserved(dtype): + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_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_gated_cuda(x, weight, torch.ones_like(x)) + expected = x * weight + bits = torch.int32 if dtype == torch.float32 else torch.int16 + assert torch.equal(actual.view(bits), expected.view(bits)) + + +@requires_cuda_gated +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +def test_cuda_gated_dgate_keeps_signed_zero_weight(dtype): + """dgate scales by the weight too, so it must not turn a -0.0 weight into +0.0. + + dgate = dy * x * rstd * w * act'(gate); with dy, rstd, act'(1) > 0 its sign is + the sign of ``x * w``, zeros included. + """ + from rl_engine.backends.cuda.norm.rmsnorm import rmsnorm_gated_cuda + + x = torch.ones(2, 128, device="cuda", dtype=dtype) + x[1].neg_() + weight = torch.full((128,), -0.0, device="cuda", dtype=dtype) + gate = torch.ones_like(x).requires_grad_() + rmsnorm_gated_cuda(x, weight, gate).backward(torch.ones_like(x)) + expected = x * weight + bits = torch.int32 if dtype == torch.float32 else torch.int16 + assert torch.equal(gate.grad.view(bits), expected.view(bits)) diff --git a/tests/models/qwen3_next/test_qwen3_next_workload.py b/tests/models/qwen3_next/test_qwen3_next_workload.py new file mode 100644 index 000000000..9710f0151 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_workload.py @@ -0,0 +1,66 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The norm workload cannot stand in for Dense or full-checkpoint evidence.""" + +import copy +from pathlib import Path + +import pytest + +from rl_engine.config.workload import WorkloadError, load_manifest, manifest_identity_hash +from rl_engine.validation.models.qwen3_next_workload import validate_norm_manifest +from rl_engine.validation.operators.gradient_adapters import get_adapter, resolve_profile_candidate + +MANIFEST = ( + Path(__file__).resolve().parents[3] + / "rl_engine/validation/models/qwen3_next_norm_manifest.json" +) + + +def test_norm_manifest_pins_real_architecture_and_separate_scope(): + manifest = load_manifest(MANIFEST) + assert manifest.model_identity["config_fingerprint"]["hidden_size"] == 2048 + assert manifest.raw["full_model_evidence"] is False + for name, gradients in ( + ("qwen3_next_rms_norm", ("dx", "dweight")), + ("rms_norm_gated", ("dx", "dweight", "dgate")), + ): + adapter = get_adapter(name) + assert tuple(t.name for t in adapter.tensors) == gradients + assert resolve_profile_candidate(adapter, "cuda_bf16", manifest)["status"] == "declared" + assert ( + resolve_profile_candidate(adapter, "triton_cuda_bf16", manifest)["status"] + == "missing_required" + ) + assert ( + resolve_profile_candidate(adapter, "cuda_bf16", load_manifest())["status"] + == "absent_not_required" + ) + + +@pytest.mark.parametrize("fault", ["architecture", "full_model", "revision", "shape", "binding"]) +def test_norm_manifest_rejects_false_evidence_even_with_regenerated_hash(fault): + raw = copy.deepcopy(load_manifest(MANIFEST).raw) + if fault == "architecture": + raw["model_identity"]["config_fingerprint"]["hidden_size"] = 4096 + elif fault == "full_model": + raw["full_model_evidence"] = True + elif fault == "revision": + raw["model_identity"]["revision"] = "main" + elif fault == "shape": + raw["representative_cases"][0]["hidden"] = 64 + else: + raw["representative_cases"][0]["fixture_id"] = "wrong_fixture" + raw["fixture_identity_sha256"] = manifest_identity_hash(raw) + with pytest.raises(WorkloadError): + validate_norm_manifest(raw) + + +def test_norm_dimension_gate_rejects_shrunk_workload(): + from rl_engine.validation.models.qwen3_next_workload import validate_norm_dimensions + + raw = load_manifest(MANIFEST).raw + with pytest.raises(WorkloadError, match="hidden 2048"): + validate_norm_dimensions(raw, "qwen3_next_rms_norm", 64, 128) + validate_norm_dimensions(raw, "rms_norm_gated", 2048, 128) 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/tests/runtime/test_dispatch.py b/tests/runtime/test_dispatch.py index 3c47bbbcf..6517d3309 100644 --- a/tests/runtime/test_dispatch.py +++ b/tests/runtime/test_dispatch.py @@ -4,6 +4,7 @@ import sys from types import ModuleType +import pytest import torch import rl_engine.platforms.device as device_module @@ -258,3 +259,67 @@ def test_executor_flow(): print("\n All infrastructure tests passed!") except Exception as e: print(f"\n Test failed with error: {e}") + + +_QWEN3_NEXT_NORMS = [ + pytest.param( + "rms_norm_gated", + OpBackend.CUDA_RMS_NORM_GATED, + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + "Qwen3NextRMSNormGatedCudaOp", + "Qwen3NextRMSNormGatedOp", + id="rms_norm_gated", + ), + pytest.param( + "qwen3_next_rms_norm", + OpBackend.CUDA_QWEN3_NEXT_RMS_NORM, + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + "Qwen3NextRMSNormCudaOp", + "Qwen3NextRMSNormOp", + id="qwen3_next_rms_norm", + ), +] + + +@pytest.mark.parametrize("op_type, cuda_backend, ref_backend, cuda_cls, ref_cls", _QWEN3_NEXT_NORMS) +def test_qwen3_next_norm_priority_is_cuda_first_with_pytorch_fallback( + op_type, cuda_backend, ref_backend, cuda_cls, ref_cls +): + """Both Qwen3-Next norms have a CUDA kernel; every other platform falls back. + + No Triton, ROCm or Ascend kernel exists yet, so those platforms must resolve + to the PyTorch reference rather than to nothing -- an operator missing from a + priority map falls through to ``OpBackend.PYTORCH_NATIVE``, which is the + logprob op, not a norm. + """ + registry = KernelRegistry() + + assert registry._priority_map["cuda"][op_type] == [cuda_backend, ref_backend] + for platform in ("rocm", "musa", "cpu", "npu"): + assert registry._priority_map[platform][op_type] == [ref_backend], platform + + +@pytest.mark.skipif(torch.version.hip is not None, reason="a cuda device maps to rocm on HIP") +@pytest.mark.parametrize("op_type, cuda_backend, ref_backend, cuda_cls, ref_cls", _QWEN3_NEXT_NORMS) +def test_qwen3_next_norm_cuda_backend_absence_falls_through_to_reference( + monkeypatch, op_type, cuda_backend, ref_backend, cuda_cls, ref_cls +): + """Without the compiled symbols the CUDA backend must be skipped, not returned. + + The zero-centred op inherits its check from ``RMSNormCudaOp.__init__``. The + lookup names a CUDA device so the CUDA-first list is walked even on a + CPU-only host; resolving for the host's own platform would reach the + reference through the CPU list without ever trying the CUDA backend. + """ + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.reference.norm import qwen3_next_rms_norm as reference + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) + monkeypatch.setattr(cuda_rmsnorm, "_C", None) + with pytest.raises(RuntimeError, match="requires the compiled rl_engine._C extension"): + getattr(cuda_rmsnorm, cuda_cls)() + + registry = KernelRegistry() + resolved = registry.get_op(op_type, device="cuda") + assert type(resolved) is getattr(reference, ref_cls) + assert cuda_backend.name in registry._failed_backends diff --git a/tests/validation/operators/test_ws1_ascend_closeout.py b/tests/validation/operators/test_ws1_ascend_closeout.py index 66c8ae4a6..b9b0fe331 100644 --- a/tests/validation/operators/test_ws1_ascend_closeout.py +++ b/tests/validation/operators/test_ws1_ascend_closeout.py @@ -187,9 +187,15 @@ def test_c2_ascend_cases_pin_real_ascend_kernels_and_sources(): def test_c3_c4_every_required_adapter_resolves_an_ascend_candidate(): manifest = load_manifest() + model_id = manifest.model_identity["model_id"] for name, adapter in GRADIENT_ADAPTERS.items(): if adapter.requirement not in ("required",): continue + # An adapter scoped to another model (the CUDA-only Qwen3-Next norms) + # resolves to `absent_not_required` for this manifest by design; this + # test covers the chain of the model whose manifest it loads. + if adapter.model_id not in (None, model_id): + continue resolved = resolve_profile_candidate(adapter, PROFILE, manifest) assert resolved["status"] == "declared", name assert resolved["expected_backend_id"] == "ascend", name @@ -197,6 +203,22 @@ def test_c3_c4_every_required_adapter_resolves_an_ascend_candidate(): assert candidate_family(str(resolved["expected_backend_id"])) == "ascend" +def test_model_scoped_adapters_are_absent_not_required_for_other_models(): + """The scoping the test above relies on must actually hold. + + Without this, skipping scoped adapters could hide one that silently resolves + to a real candidate for the wrong model. + """ + manifest = load_manifest() + model_id = manifest.model_identity["model_id"] + scoped = [a for a in GRADIENT_ADAPTERS.values() if a.model_id not in (None, model_id)] + assert scoped, "expected at least the Qwen3-Next norm adapters to be model-scoped" + for adapter in scoped: + resolved = resolve_profile_candidate(adapter, PROFILE, manifest) + assert resolved["status"] == "absent_not_required", adapter.op_name + assert resolved["candidate_path"] is None, adapter.op_name + + def test_c4_adapter_status_matrix_has_no_red_ascend_rows(): rows = [r for r in gradient_adapter_status_matrix() if r.backend_profile == PROFILE] assert rows diff --git a/tests/validation/operators/test_ws1_gtest_gpu.py b/tests/validation/operators/test_ws1_gtest_gpu.py index bcaebc2f5..60ac5102c 100644 --- a/tests/validation/operators/test_ws1_gtest_gpu.py +++ b/tests/validation/operators/test_ws1_gtest_gpu.py @@ -36,6 +36,8 @@ def test_all_ws1_single_ops_are_registered(): names = set(operator_names()) assert { "rms_norm", + "rms_norm_gated", + "qwen3_next_rms_norm", "qk_norm", "det_gemm", "attention", 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..aaeba82c3 --- /dev/null +++ b/tools/validation/models/qwen3_next_norm_evidence.py @@ -0,0 +1,348 @@ +#!/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 +import torch.nn.functional as F + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.backends.cuda.norm.rmsnorm import ( # noqa: E402 + Qwen3NextRMSNormCudaOp, + Qwen3NextRMSNormGatedCudaOp, +) +from rl_engine.reference.norm.qwen3_next_rms_norm import ( # noqa: E402 + Qwen3NextRMSNormGatedOp, + Qwen3NextRMSNormOp, +) + +EPS = 1e-6 + + +# --------------------------------------------------------------------------- # +# Ops. Each candidate is fn(*row_inputs, weight) -> y, plus whether it has a +# backward. Row inputs are sliced together for the row-invariance check. +# --------------------------------------------------------------------------- # + + +def zero_centred_candidates(hidden: int) -> 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, w): + x64 = x.double() + return x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + EPS) * (1.0 + w.double()) + + +def gated_candidates(hidden: int) -> dict[str, dict[str, Any]]: + cuda_op, ref_op = Qwen3NextRMSNormGatedCudaOp(), Qwen3NextRMSNormGatedOp() + cands: dict[str, dict[str, Any]] = { + "rl-kernel CUDA": {"fn": lambda x, g, w: cuda_op(x, w, g, eps=EPS), "backward": True}, + "rl-kernel PyTorch reference": { + "fn": lambda x, g, w: ref_op(x, w, g, eps=EPS), + "backward": True, + }, + } + try: + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNormGated + + module = Qwen3NextRMSNormGated(hidden, eps=EPS).cuda() + + def hf(x, g, w, module=module): + return torch.func.functional_call(module, {"weight": w}, (x, g)) + + cands["transformers Qwen3NextRMSNormGated (cast-first)"] = {"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 RMSNormGated + + gated = RMSNormGated(hidden, eps=EPS, norm_before_gate=True).cuda() + + def vllm_fwd(x, g, w, gated=gated): + gated.weight.data = w.detach() + return gated.forward_cuda(x, g) + + cands["vLLM RMSNormGated (forward only)"] = {"fn": vllm_fwd, "backward": False} + except ImportError: + pass + return cands + + +def gated_golden(x, g, w): + """vLLM's convention (the one #468 implements) in FP64: x * rstd * w * silu(gate).""" + + x64 = x.double() + return (x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + EPS) * w.double()) * F.silu( + g.double() + ) + + +OPS: dict[str, dict[str, Any]] = { + "zero_centred_rmsnorm": { + "hidden": 2048, # decoder and final norms + "row_inputs": 1, + "weight": lambda h, gen: torch.randn(h, device="cuda", generator=gen) * 0.1, + "candidates": zero_centred_candidates, + "golden": zero_centred_golden, + "timed_rows": (1024, 4096, 16384, 65536), + }, + "gated_rmsnorm": { + "hidden": 128, # GDN value head dim; rows are tokens x heads + "row_inputs": 2, + "weight": lambda h, gen: 1.0 + torch.randn(h, device="cuda", generator=gen) * 0.1, + "candidates": gated_candidates, + "golden": gated_golden, + "timed_rows": (4096, 16384, 65536, 262144), + }, +} + + +# --------------------------------------------------------------------------- # +# Measurements +# --------------------------------------------------------------------------- # + + +def _inputs(spec, rows: int, seed: int, dtype=torch.bfloat16): + g = torch.Generator(device="cuda").manual_seed(seed) + h = spec["hidden"] + row_inputs = [(torch.randn(rows, h, device="cuda", generator=g) * 2).to(dtype)] + for _ in range(spec["row_inputs"] - 1): + row_inputs.append(torch.randn(rows, h, device="cuda", generator=g).to(dtype)) + w = spec["weight"](h, g).to(dtype) + up = torch.randn(rows, h, device="cuda", generator=g).to(dtype) + return row_inputs, w, up + + +def _grads(fn, row_inputs, w, up, dtype=None): + def leaf(t): + return (t if dtype is None else t.to(dtype)).detach().clone().requires_grad_(True) + + rl, wl = [leaf(t) for t in row_inputs], leaf(w) + out = fn(*rl, wl) + out.backward(up if dtype is None else up.to(dtype)) + return out.detach(), [t.grad for t in rl], wl.grad + + +def accuracy(spec, cands, rows: int, seed: int) -> dict[str, Any]: + row_inputs, w, up = _inputs(spec, rows, seed) + ref_out, ref_drows, ref_dw = _grads(spec["golden"], row_inputs, w, up, torch.float64) + result = {} + for name, c in cands.items(): + entry: dict[str, Any] = {} + if c["backward"]: + out, drows, dw = _grads(c["fn"], row_inputs, w, up) + pairs = [("dx", drows[0], ref_drows[0]), ("dweight", dw, ref_dw)] + if len(drows) > 1: + pairs.append(("dgate", drows[1], ref_drows[1])) + for key, got, ref in pairs: + 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"](*row_inputs, 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(spec, cands, seeds=(3, 4, 5), rows: int = 4096, step: int = 16): + """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: + row_inputs, w, up = _inputs(spec, rows, seed) + if c["backward"]: + full_out, full_drows, _ = _grads(c["fn"], row_inputs, w, up) + else: + with torch.no_grad(): + full_out = c["fn"](*row_inputs, w) + for i in range(0, rows, step): + part = [t[i : i + 1] for t in row_inputs] + if c["backward"]: + out, drows, _ = _grads(c["fn"], part, w, up[i : i + 1]) + dx_bad += not all(torch.equal(d[0], fd[i]) for d, fd in zip(drows, full_drows)) + else: + with torch.no_grad(): + out = c["fn"](*part, 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(spec, cands) -> dict[str, Any]: + result: dict[str, Any] = {name: {} for name in cands} + for rows in spec["timed_rows"]: + row_inputs, w, up = _inputs(spec, 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(*row_inputs, w)) + if c["backward"]: + leaves = [t.detach().clone().requires_grad_(True) for t in [*row_inputs, w]] + out = c["fn"](*leaves) + row["backward_us"] = _time_us( + lambda o=out, lv=leaves: torch.autograd.grad(o, lv, 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, spec) -> dict[str, Any]: + cands = spec["candidates"](spec["hidden"]) + print(f"[{name}] candidates: {', '.join(cands)}", flush=True) + report = { + "hidden": spec["hidden"], + "accuracy": {str(r): accuracy(spec, cands, r, seed=r) for r in (257, 4096)}, + "row_invariance": row_invariance(spec, cands), + "latency": latency(spec, cands), + } + for cand, entry in report["row_invariance"].items(): + print(f" BI {cand}: {entry}", flush=True) + return report + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--ops", default=",".join(OPS), help="comma list of ops") + args = parser.parse_args() + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + 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(), + "eps": EPS, + "dtype": "bfloat16", + "ops": {name: run_op(name, OPS[name]) for name in args.ops.split(",")}, + } + 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/ws1_gdn_provider_agreement.py b/tools/validation/models/ws1_gdn_provider_agreement.py new file mode 100644 index 000000000..5b2ee3567 --- /dev/null +++ b/tools/validation/models/ws1_gdn_provider_agreement.py @@ -0,0 +1,496 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Measure how closely the GDN decode-step goldens track vLLM's providers (RFC #428 C6). + +``tests/models/qwen3_next/check_gdn_recurrent_golden.py`` asserts regression bounds; this runner +prints the measured values behind them, so the numbers quoted in +``docs/archive/design/rfc428-c6-gdn-recurrent-replay.md`` can be reproduced. Inputs and seeds +are the check file's own helpers, imported from it rather than copied. + +Measurements, emitted as one JSON document on stdout: + +* ``recurrent``: provider vs golden for the packed recurrent decode, per + (batch, state dtype): max|d out|, max|d state| and bitwise mismatch counts. +* ``conv``: provider vs golden for ``causal_conv1d_update``, per (batch, cache dtype): + output mismatch count and max|diff|, and whether the rolled state is bitwise equal. +* ``conv`` again with Triton FP fusion disabled, plus ``conv_kernels``: per compiled + variant of the provider's conv-update kernel, its ``enable_fp_fusion`` option and + instruction counts from the PTX and from the SASS. +* ``conv_silu`` localises the fp32-cache conv mismatches, on the same inputs widened to + fp32 with an fp32 output: the pre-activation comparison (activation off), the SiLU + comparison with its ULP histogram, and four Triton SiLU variants applied to the + golden's pre-activation values, each compared bitwise with both sides. +* ``conv_noact_bf16``: the no-activation, bf16-output specialization against the golden, + whether it equals the fp32-output run rounded to bf16, each mismatch's size, and the + compiled variants of both specializations; with fusion on and off. + +The fusion-off arm runs in a child process with ``TRITON_DEFAULT_FP_FUSION=0`` and a +private ``TRITON_CACHE_DIR``. Flipping ``triton.knobs.language.default_fp_fusion`` +in-process does not work: Triton's in-memory kernel cache is keyed on the launch +kwargs, and the knob is read only after a cache miss, so the fused variant is reused. + +Count FMAs in the SASS, not only the PTX. With fusion on, Triton emits plain +``mul.f32``/``add.f32`` and leaves ptxas free to contract them into ``FFMA``; with +fusion off it passes ``--fmad=false`` to ptxas. A PTX with no ``fma.rn.f32`` can still +run FMAs. ``fusion_check`` reports whether each arm compiled what it claims to. + +Requires CUDA and vLLM 0.30.0. Run it from a clean checkout so ``git_dirty`` is false, +and write the result outside the checkout: an untracked file inside it would itself +make the next run report ``git_dirty: true``. + + python tools/validation/models/ws1_gdn_provider_agreement.py \\ + > "${TMPDIR:-/tmp}/gdn_provider_agreement.json" +""" + +from __future__ import annotations + +import argparse +import importlib.util +import json +import math +import os +import platform +import re +import subprocess +import sys +import tempfile +from pathlib import Path +from typing import Any + +import torch + +REPO_ROOT = Path(__file__).resolve().parents[3] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +_CHECK_FILE = REPO_ROOT / "tests" / "check_gdn_recurrent_golden.py" +_STATE_DTYPES = {"fp32": torch.float32, "bf16": torch.bfloat16} +_INT_VIEW = {2: torch.int16, 4: torch.int32} +# FLA_USE_FAST_OPS swaps the kernel's exp/log for fast_expf/fast_logf; +# TRITON_DEFAULT_FP_FUSION decides whether ptxas may contract mul/add. +_PROVENANCE_ENV = ("FLA_USE_FAST_OPS", "TRITON_DEFAULT_FP_FUSION") +# Opcode patterns. PTX: a rounding-qualified op (".rn") may not be contracted by ptxas; +# the plain form may. SASS: count opcode tokens, including modifiers such as FFMA.FTZ. +_PTX_OPS = { + "fma_rn_f32": r"\bfma\.rn\.f32\b", + "mul_f32": r"\bmul\.f32\b", + "add_f32": r"\badd\.f32\b", + "mul_rn_f32": r"\bmul\.rn\.f32\b", + "add_rn_f32": r"\badd\.rn\.f32\b", + # Packed pairs (sm_100), approximate exp and division. + "fma_rn_f32x2": r"\bfma\.rn\.f32x2\b", + "mul_f32x2": r"\bmul(\.rn)?\.f32x2\b", + "add_f32x2": r"\badd(\.rn)?\.f32x2\b", + "ex2_approx": r"\bex2\.approx", + "div_full_f32": r"\bdiv\.full\.f32\b", +} +_SASS_OPS = {"FFMA": r"\bFFMA[\w.]*", "FMUL": r"\bFMUL[\w.]*", "FADD": r"\bFADD[\w.]*"} +# Triton SiLU variants for conv_silu, by MODE of _triton_silu's kernel. +_SILU_VARIANTS = ("div_exp", "divrn_exp", "div_libexp", "divrn_libexp") + + +def _load_check_module(): + """Import the check file by path; ``tests/`` is not a package.""" + spec = importlib.util.spec_from_file_location("check_gdn_recurrent_golden", _CHECK_FILE) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _bit_mismatches(a: torch.Tensor, b: torch.Tensor) -> int: + """Elements whose bit patterns differ (NaN-safe, unlike ``a != b``).""" + assert a.dtype == b.dtype and a.shape == b.shape, (a.dtype, b.dtype, a.shape, b.shape) + view = _INT_VIEW[a.element_size()] + return int((a.contiguous().view(view) != b.contiguous().view(view)).sum()) + + +def _max_abs_diff(a: torch.Tensor, b: torch.Tensor) -> float: + return (a.float() - b.float()).abs().max().item() + + +def _git(*args: str) -> str | None: + """Stripped stdout, or ``None`` if git is missing or the command fails.""" + try: + return subprocess.run( + ["git", *args], cwd=REPO_ROOT, capture_output=True, text=True, check=True + ).stdout.strip() + except (OSError, subprocess.CalledProcessError): + return None + + +def _provenance() -> dict[str, Any]: + import triton + import vllm + + import rl_engine + + status = _git("status", "--porcelain") + return { + # None means "unknown" (no git, or not a checkout), never "clean". + "git_commit": _git("rev-parse", "HEAD"), + "git_dirty": None if status is None else bool(status), + # Which tree was imported, and the env knobs that change the provider's code. + "rl_engine_file": rl_engine.__file__, + "env": {key: os.environ.get(key) for key in _PROVENANCE_ENV}, + "python": platform.python_version(), + "torch": torch.__version__, + "triton": triton.__version__, + "vllm": vllm.__version__, + "device": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "default_fp_fusion": bool(triton.knobs.language.default_fp_fusion), + } + + +def _measure_recurrent(check, batches: list[int]) -> list[dict[str, Any]]: + bounds = {name: check._RECURRENT_BOUNDS[dtype] for name, dtype in _STATE_DTYPES.items()} + rows = [] + for batch in batches: + for name, state_dtype in _STATE_DTYPES.items(): + # Same call as test_golden_matches_packed_decode_provider. + inp = check._inputs(batch, batch + 2, state_dtype, torch.bfloat16, seed=batch) + out_ref, state_ref = check._run_provider(inp) + out_got, state_got = check._run_golden(inp) + rows.append( + { + "batch": batch, + "state_dtype": name, + "max_abs_diff_out": _max_abs_diff(out_got, out_ref), + "max_abs_diff_state": _max_abs_diff(state_got, state_ref), + "out_mismatch_elements": _bit_mismatches(out_got, out_ref), + "out_elements": out_ref.numel(), + "state_mismatch_elements": _bit_mismatches(state_got, state_ref), + "state_elements": state_ref.numel(), + "asserted_out_atol": bounds[name][0], + "asserted_state_atol": bounds[name][1], + } + ) + return rows + + +def _measure_conv(check, batches: list[int]) -> list[dict[str, Any]]: + import triton + + fp_fusion = bool(triton.knobs.language.default_fp_fusion) # this process's setting + rows = [] + for batch in batches: + for name, cache_dtype in _STATE_DTYPES.items(): + # Same call as the test_conv_* provider comparisons. + inp = check._conv_inputs(batch, cache_dtype, seed=batch) + (out_ref, state_ref), (out_got, state_got) = check._run_conv_pair(inp) + rows.append( + { + "batch": batch, + "cache_dtype": name, + "fp_fusion": fp_fusion, + "out_mismatch_elements": _bit_mismatches(out_got, out_ref), + "out_elements": out_ref.numel(), + "max_abs_diff_out": _max_abs_diff(out_got, out_ref), + "state_bitwise_equal": torch.equal(state_got, state_ref), + } + ) + return rows + + +def _count(patterns: dict[str, str], text: str) -> dict[str, int]: + return {name: len(re.findall(pattern, text)) for name, pattern in patterns.items()} + + +def _compiled_conv_variants() -> dict[tuple[str, Any], Any]: + """Every compiled variant of vLLM's conv-update kernel in this process, by cache key.""" + from vllm.model_executor.layers.mamba.ops import causal_conv1d as conv_module + + kernel = conv_module._causal_conv1d_update_kernel + while not hasattr(kernel, "device_caches") and hasattr(kernel, "fn"): + kernel = kernel.fn + return { + (str(device), key): compiled + for device, cache in kernel.device_caches.items() + for key, compiled in cache[0].items() + } + + +def _conv_kernel_variants(keys: Any = None) -> list[dict[str, Any]] | dict[str, str]: + """Options and instruction counts per compiled variant (only ``keys``, if given).""" + try: + variants = [] + for (device, key), compiled in _compiled_conv_variants().items(): + if keys is not None and (device, key) not in keys: + continue + row: dict[str, Any] = { + "device": device, + "enable_fp_fusion": getattr(compiled.metadata, "enable_fp_fusion", None), + "ptx": _count(_PTX_OPS, compiled.asm["ptx"]), + } + try: + row["sass"] = _count(_SASS_OPS, compiled.asm["sass"]) + except Exception as exc: # needs cuobjdump; report, do not fail + row["sass"] = {"unavailable": f"{type(exc).__name__}: {exc}"} + variants.append(row) + return variants + except Exception as exc: # introspection of Triton internals; report, do not fail + return {"unavailable": f"{type(exc).__name__}: {exc}"} + + +def _new_variant_keys(before: set[Any]) -> set[Any] | None: + try: + return set(_compiled_conv_variants()) - before + except Exception: # introspection of Triton internals; report, do not fail + return None + + +def _bf16_ulps(diff: float, ref: float) -> float | None: + """``diff`` in units of one bf16 ULP at the magnitude of ``ref`` (None at zero).""" + if ref == 0.0: + return None if diff else 0.0 + return diff / 2.0 ** (math.floor(math.log2(abs(ref))) - 7) + + +def _conv_noact_bf16(check, batches: list[int]) -> dict[str, Any]: + """The no-activation conv with a bf16 output, against its fp32-output twin. + + The provider casts x to the fp32 cache dtype before launching, so the two runs + compute the same thing and differ only in the output dtype, i.e. in which Triton + specialization runs. Each specialization runs in its own loop so that the compiled + variants it adds can be told apart. + """ + import triton + + fp_fusion = bool(triton.knobs.language.default_fp_fusion) + try: + before = set(_compiled_conv_variants()) + except Exception: # introspection of Triton internals; report, do not fail + before = set() + bf16_out = {} + for batch in batches: + inp = check._conv_inputs(batch, torch.float32, seed=batch) + bf16_out[batch] = check._run_conv_pair(inp, activation=None) + bf16_keys = _new_variant_keys(before) + fp32_out = {} + for batch in batches: + inp = check._conv_inputs(batch, torch.float32, seed=batch) + fp32_out[batch] = check._run_conv_pair(dict(inp, x=inp["x"].float()), activation=None) + fp32_keys = _new_variant_keys(before | (bf16_keys or set())) + + rows = [] + for batch in batches: + (b_ref, _), (b_got, _) = bf16_out[batch] + (d_ref, _), (d_got, _) = fp32_out[batch] + differ = torch.nonzero(b_ref.view(torch.int16) != b_got.view(torch.int16)).tolist() + mismatches = [] + for r, ch in differ: + provider, golden = b_ref[r, ch].item(), b_got[r, ch].item() + mismatches.append( + { + "abs_out": abs(golden), + "abs_diff": abs(provider - golden), + "bf16_ulps": _bf16_ulps(abs(provider - golden), golden), + } + ) + rows.append( + { + "batch": batch, + "fp_fusion": fp_fusion, + "out_mismatch_elements": len(differ), + "out_elements": b_ref.numel(), + "provider_eq_rne_of_fp32_out": torch.equal(b_ref, d_ref.to(torch.bfloat16)), + "fp32_out_mismatch_elements": _bit_mismatches(d_got, d_ref), + "mismatches": mismatches, + } + ) + return { + "batches": rows, + "bf16_out_kernels": _conv_kernel_variants(bf16_keys) if bf16_keys is not None else None, + "fp32_out_kernels": _conv_kernel_variants(fp32_keys) if fp32_keys is not None else None, + } + + +def _triton_silu(): + """``apply(x, mode)``: one of the four SiLU formulations, evaluated by Triton.""" + import triton + import triton.language as tl + from triton.language.extra import libdevice + + @triton.jit + def silu(x_ptr, y_ptr, n, MODE: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask, other=0.0) + if MODE == 0: # the provider's source: x / (1 + tl.exp(-x)) + y = x / (1 + tl.exp(-x)) + elif MODE == 1: + y = tl.math.div_rn(x, 1 + tl.exp(-x)) + elif MODE == 2: + y = x / (1 + libdevice.exp(-x)) + else: + y = tl.math.div_rn(x, 1 + libdevice.exp(-x)) + tl.store(y_ptr + offs, y, mask=mask) + + def apply(x: torch.Tensor, mode: int) -> torch.Tensor: + flat = x.contiguous().view(-1) + y = torch.empty_like(flat) + silu[(triton.cdiv(flat.numel(), 1024),)](flat, y, flat.numel(), MODE=mode, BLOCK=1024) + return y.view_as(x) + + return apply + + +def _ulp_histogram(a: torch.Tensor, b: torch.Tensor) -> dict[str, int]: + """Differing fp32 elements by bit-pattern distance.""" + d = (a.contiguous().view(torch.int32).long() - b.contiguous().view(torch.int32).long()).abs() + d = d[d != 0] + return { + "1": int((d == 1).sum()), + "2": int((d == 2).sum()), + "3-4": int(((d == 3) | (d == 4)).sum()), + ">4": int((d > 4).sum()), + } + + +def _conv_silu(check, batches: list[int]) -> list[dict[str, Any]]: + """Where the fp32-cache conv mismatches come from, on fp32 inputs and outputs.""" + silu = _triton_silu() + rows = [] + for batch in batches: + inp = check._conv_inputs(batch, torch.float32, seed=batch) + inp32 = dict(inp, x=inp["x"].float()) # the same values, kept in fp32 throughout + (pre_ref, _), (pre_got, _) = check._run_conv_pair(inp32, activation=None) + (out_ref, _), (out_got, _) = check._run_conv_pair(inp32) + variants = {name: silu(pre_got, mode) for mode, name in enumerate(_SILU_VARIANTS)} + rows.append( + { + "batch": batch, + "elements": out_ref.numel(), + "preactivation_mismatch_elements": _bit_mismatches(pre_got, pre_ref), + "silu_mismatch_elements": _bit_mismatches(out_got, out_ref), + "silu_ulp_histogram": _ulp_histogram(out_got, out_ref), + # Applied to the golden's pre-activation values. + "triton_silu_vs_provider": { + k: _bit_mismatches(v, out_ref) for k, v in variants.items() + }, + "triton_silu_vs_golden": { + k: _bit_mismatches(v, out_got) for k, v in variants.items() + }, + } + ) + return rows + + +def _conv_report(check, batches: list[int]) -> dict[str, Any]: + conv = _measure_conv(check, batches) + kernels = _conv_kernel_variants() # before the no-activation runs add their own + return { + "conv": conv, + "conv_kernels": kernels, + "conv_noact_bf16": _conv_noact_bf16(check, batches), + } + + +def _arm_variants(arm: dict[str, Any]) -> list[Any] | None: + """Every compiled variant an arm reports, or None if any listing is unavailable.""" + noact = arm["conv_noact_bf16"] + lists = [arm["conv_kernels"], noact["bf16_out_kernels"], noact["fp32_out_kernels"]] + if not all(isinstance(x, list) for x in lists): + return None + return [v for x in lists for v in x] + + +def _conv_report_without_fusion(batches: str) -> dict[str, Any]: + """Rerun the conv arm in a child process that compiles with fusion off.""" + with tempfile.TemporaryDirectory(prefix="triton-nofusion-") as cache_dir: + env = dict(os.environ, TRITON_DEFAULT_FP_FUSION="0", TRITON_CACHE_DIR=cache_dir) + child = subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--conv-only", "--batches", batches], + env=env, + stdout=subprocess.PIPE, + text=True, + check=True, + ) + return json.loads(child.stdout) + + +def _fusion_check(on: Any, off: Any) -> dict[str, Any]: + """Whether each arm's compiled variants carry the fusion setting it claims.""" + + def flags(variants: Any) -> list[Any] | None: + if not isinstance(variants, list): + return None + return [v["enable_fp_fusion"] for v in variants] + + on_flags, off_flags = flags(on), flags(off) + return { + "fusion_on_variants": on_flags, + "fusion_off_variants": off_flags, + "fusion_off_effective": bool(off_flags) and all(f is False for f in off_flags), + } + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument( + "--batches", + default="1,4,17,64", + help="comma-separated batch sizes (default: the check file's 1,4,17,64)", + ) + parser.add_argument( + "--skip-fusion-off", + action="store_true", + help="do not rerun the conv comparison with Triton FP fusion disabled", + ) + # Internal: the fusion-off child process prints only the conv arm. + parser.add_argument("--conv-only", action="store_true", help=argparse.SUPPRESS) + return parser.parse_args() + + +def main() -> int: + args = parse_args() + if not torch.cuda.is_available(): + print("CUDA is required", file=sys.stderr) + return 2 + try: + import vllm # noqa: F401 + except ImportError: + print("vLLM is required", file=sys.stderr) + return 2 + + batches = [int(b) for b in args.batches.split(",") if b.strip()] + check = _load_check_module() + if args.conv_only: + print(json.dumps(_conv_report(check, batches), sort_keys=True)) + return 0 + + report: dict[str, Any] = { + "runner": "tools/validation/models/ws1_gdn_provider_agreement.py", + "provenance": _provenance(), + "inputs": { + "source": ( + "tests/models/qwen3_next/check_gdn_recurrent_golden.py " "(_inputs, _conv_inputs)" + ), + "seed": "seed=batch for every case", + "batches": batches, + "recurrent": "bf16 I/O, use_qk_l2norm_in_kernel=True, num_blocks=batch+2", + "conv": "bias=True, activation=silu, dim_first=True, W=4, dim=8192", + "conv_silu": "the conv inputs with x widened to fp32; fp32 cache and output", + "conv_noact_bf16": "the conv inputs with activation=None; fp32 cache", + }, + "recurrent": _measure_recurrent(check, batches), + } + fused = _conv_report(check, batches) + report["conv"] = fused["conv"] + report["conv_kernels"] = {"fusion_on": fused["conv_kernels"]} + report["conv_noact_bf16"] = {"fusion_on": fused["conv_noact_bf16"]} + report["conv_silu"] = _conv_silu(check, batches) + if not args.skip_fusion_off: + unfused = _conv_report_without_fusion(args.batches) + report["conv"] += unfused["conv"] + report["conv_kernels"]["fusion_off"] = unfused["conv_kernels"] + report["conv_noact_bf16"]["fusion_off"] = unfused["conv_noact_bf16"] + report["fusion_check"] = _fusion_check(_arm_variants(fused), _arm_variants(unfused)) + print(json.dumps(report, indent=2, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/validation/operators/check_forward_invariance.py b/tools/validation/operators/check_forward_invariance.py index c7a836f2c..6830d8315 100644 --- a/tools/validation/operators/check_forward_invariance.py +++ b/tools/validation/operators/check_forward_invariance.py @@ -18,8 +18,9 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) -from rl_engine.config.workload import load_manifest # noqa: E402 +from rl_engine.config.workload import load_manifest, workload_report # noqa: E402 from rl_engine.contracts.numerical import resolve_dtype_policy # noqa: E402 +from rl_engine.validation.models.qwen3_next_workload import validate_norm_dimensions # noqa: E402 from rl_engine.validation.operators import ( # noqa: E402 BackendProvenance, assert_forward_batch_invariant, @@ -113,6 +114,12 @@ def parse_args() -> argparse.Namespace: if adapter.requirement != "absent_not_required" ] parser = argparse.ArgumentParser(description="WS1 C3 forward invariance GPU gate") + parser.add_argument( + "--manifest", + type=pathlib.Path, + default=None, + help="Workload manifest JSON (default: the Qwen3-8B Dense C2 manifest)", + ) parser.add_argument("--op", choices=sorted(runnable), default="rms_norm") parser.add_argument( "--candidate", required=True, help="Manifest-declared CUDA/Triton/Ascend candidate" @@ -146,7 +153,8 @@ def main() -> None: raise SystemExit("ERROR: --vocab must cover every fixed C2 workload token id") contract = load_contract() - manifest = load_manifest() + manifest = load_manifest(args.manifest) + validate_norm_dimensions(manifest.raw, args.op, args.hidden, args.head_dim) adapter = get_adapter(args.op) if adapter.requirement == "layout_supported": raise SystemExit( @@ -233,7 +241,9 @@ def main() -> None: ) if args.json: - print(json.dumps(report.to_dict(), indent=2, default=str)) + payload = report.to_dict() + payload["workload"] = workload_report(manifest) + print(json.dumps(payload, indent=2, default=str)) else: _summarize(report) if not report.passed: diff --git a/tools/validation/operators/check_gradient_invariance.py b/tools/validation/operators/check_gradient_invariance.py index db81d8c58..178a5a4a6 100644 --- a/tools/validation/operators/check_gradient_invariance.py +++ b/tools/validation/operators/check_gradient_invariance.py @@ -18,8 +18,9 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) -from rl_engine.config.workload import load_manifest # noqa: E402 +from rl_engine.config.workload import load_manifest, workload_report # noqa: E402 from rl_engine.contracts.numerical import resolve_dtype_policy # noqa: E402 +from rl_engine.validation.models.qwen3_next_workload import validate_norm_dimensions # noqa: E402 from rl_engine.validation.operators import ( # noqa: E402 BackendProvenance, assert_gradient_batch_invariant, @@ -125,6 +126,12 @@ def parse_args() -> argparse.Namespace: if adapter.requirement != "absent_not_required" ] parser = argparse.ArgumentParser(description="WS1 C4 gradient invariance GPU gate") + parser.add_argument( + "--manifest", + type=pathlib.Path, + default=None, + help="Workload manifest JSON (default: the Qwen3-8B Dense C2 manifest)", + ) parser.add_argument("--op", choices=sorted(runnable), default="rms_norm") parser.add_argument( "--candidate", required=True, help="Manifest-declared CUDA/Triton/Ascend candidate" @@ -159,7 +166,8 @@ def main() -> None: raise SystemExit(f"ERROR: C4 required-profile evidence needs a real device: {exc}") from exc contract = load_contract() - manifest = load_manifest() + manifest = load_manifest(args.manifest) + validate_norm_dimensions(manifest.raw, args.op, args.hidden, args.head_dim) adapter = get_adapter(args.op) if adapter.requirement == "layout_supported": # Pack is the same PyTorch layout op under both profiles and is not a C2 @@ -255,7 +263,9 @@ def main() -> None: ) from exc if args.json: - print(json.dumps(report.to_dict(), indent=2, default=str)) + payload = report.to_dict() + payload["workload"] = workload_report(manifest) + print(json.dumps(payload, indent=2, default=str)) else: _summarize(report) if not report.passed: