diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4634569cd..ed3c3fd72 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 tests/validation/common/test_tensor_identity.py tests/models/qwen3_next/test_qwen3_next_forward_contract.py tests/models/qwen3_next/test_qwen3_next_tp_blocks.py tests/models/qwen3_next/test_qwen3_next_tp_gdn.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..071d42a60 --- /dev/null +++ b/.github/workflows/qwen3-next-provider-gpu.yml @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 +# RFC #428 Qwen3-Next provider comparisons: the C1/C3/C4 norm providers, the C6 +# GDN decode-step goldens, the shared GEMM/MoE route-combine provider and the TP4 +# GDN/convolution/attention mixers, 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/**" + - "csrc/cuda/attention/**" + - "rl_engine/models/qwen3_next/**" + - "rl_engine/integrations/engines/train/vllm/qwen3_next*" + - "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 + # The FP32 router needs IEEE Triton dots from process start. + env: + TRITON_F32_DEFAULT: ieee + run: | + python3 -m pytest -q tests/models/qwen3_next/check_qwen3_next_norm_providers.py tests/models/qwen3_next/check_gdn_recurrent_golden.py \ + tests/models/qwen3_next/check_qwen3_next_forward.py tests/models/qwen3_next/check_qwen3_next_attention.py \ + tests/models/qwen3_next/check_qwen3_next_gdn_bridge.py tests/models/qwen3_next/check_qwen3_next_gdn_sequence.py \ + tests/models/qwen3_next/check_qwen3_next_conv_bridge.py tests/models/qwen3_next/check_qwen3_next_core_matrix.py \ + tests/models/qwen3_next/check_qwen3_next_shared_core.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/attention/deterministic_attention.cu b/csrc/cuda/attention/deterministic_attention.cu index 82ec824db..aeae9231c 100644 --- a/csrc/cuda/attention/deterministic_attention.cu +++ b/csrc/cuda/attention/deterministic_attention.cu @@ -17,7 +17,6 @@ namespace { -constexpr int64_t kDeterministicAttentionHeadDim = 128; constexpr int kSoftmaxThreads = 256; // --------------------------------------------------------------------------- @@ -220,15 +219,18 @@ void check_deterministic_attention_inputs( const int64_t Hkv = k.size(1); const int64_t Skv = k.size(2); - TORCH_CHECK(D == kDeterministicAttentionHeadDim, - "deterministic_attention: head dim D must be ", - kDeterministicAttentionHeadDim, - ", got ", - D); +#if defined(USE_ROCM) + TORCH_CHECK(D == 128, "deterministic_attention: ROCm head dim D must be 128, got ", D); +#else + TORCH_CHECK(D == 128 || D == 256, + "deterministic_attention: CUDA head dim D must be 128 or 256, got ", D); +#endif TORCH_CHECK(k.size(0) == B && v.size(0) == B, "deterministic_attention: batch size mismatch between q/k/v"); TORCH_CHECK(v.size(1) == Hkv && v.size(2) == Skv && k.size(3) == D && v.size(3) == D, "deterministic_attention: k/v shape mismatch"); + TORCH_CHECK(B > 0 && Hq > 0 && Hkv > 0, + "deterministic_attention: batch size and head counts must be positive"); TORCH_CHECK(Hq % Hkv == 0, "deterministic_attention: Hq (", Hq, @@ -573,6 +575,17 @@ std::vector deterministic_attention_backward( double scale, torch::optional key_padding_mask) { + check_deterministic_attention_inputs(q, k, v, key_padding_mask); + TORCH_CHECK(grad_output.device() == q.device() && grad_output.scalar_type() == q.scalar_type(), + "deterministic_attention_backward: grad_output must match q device and dtype"); + TORCH_CHECK(grad_output.sizes() == q.sizes(), + "deterministic_attention_backward: grad_output must have q shape"); + TORCH_CHECK(P.device() == q.device() && P.scalar_type() == at::kFloat, + "deterministic_attention_backward: P must be FP32 on the input device"); + TORCH_CHECK(P.dim() == 4 && P.size(0) == q.size(0) && P.size(1) == q.size(1) && + P.size(2) == q.size(2) && P.size(3) == k.size(2), + "deterministic_attention_backward: P must have shape [B, Hq, Sq, Skv]"); + const at::cuda::OptionalCUDAGuard device_guard(at::device_of(q)); auto dO = grad_output.contiguous(); 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..be74cb6d0 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -32,4 +32,7 @@ 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) +- [Qwen3-Next MoE route / combine](qwen3-next-moe-route-combine.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/attention.md b/docs/operators/attention.md index 52f8517ea..be71c7e7e 100644 --- a/docs/operators/attention.md +++ b/docs/operators/attention.md @@ -265,7 +265,7 @@ for measured peak memory at representative shapes. ## Known Limitations -- First version: `D=128` only (Qwen3-8B alignment). +- Head dim: `D=128` (Qwen3-8B) on CUDA and ROCm; `D=256` (Qwen3-Next full attention) on CUDA only. - Supported dtypes: BF16, FP16. - Full materialization of scores/P limits practical sequence length. - `Hq` must be divisible by `Hkv` (raises `ValueError` otherwise). diff --git a/docs/operators/qwen3-next-moe-route-combine.md b/docs/operators/qwen3-next-moe-route-combine.md new file mode 100644 index 000000000..07804b8bc --- /dev/null +++ b/docs/operators/qwen3-next-moe-route-combine.md @@ -0,0 +1,146 @@ +# Qwen3-Next MoE route / combine contract + +## Summary + +Qwen3-Next's sparse MoE block routes every token to 10 of 512 experts (expert +width 512, TP4-local width 128) and adds a gated shared expert. RFC #428 row +`moe_route_combine_contract` asks for a routed MoE whose output for a token is a +pure function of that token: the same bits whatever else is in the batch, in +training replay and in rollout. This page states the contract and how +`rl_engine.models.qwen3_next.qwen3_next_forward.shared_moe` meets it. + +Upstream references: `transformers` `Qwen3NextSparseMoeBlock` (5.17.0) and vLLM +0.30.0 `fused_topk` + `fused_experts`. + +## Entry Point + +```python +from rl_engine.models.qwen3_next.qwen3_next_forward import shared_moe, stable_top10_routes +from rl_engine.models.qwen3_next.qwen3_next_tp_blocks import TP4MoE + +routed, routes = shared_moe(x, router_weight, gate_up, down) # TP-local, not reduced +block = TP4MoE(group=tp_group, device="cuda") # routed + shared, reduced +block.load_hf_weights(hf_layer_mlp_weights) +y = block(hidden) +``` + +`TRITON_F32_DEFAULT=ieee` must be set before the process starts: the FP32 router +GEMM refuses to run on TF32 Triton dots. + +## Contract + +| Stage | Order and dtype | Atomics | +| --- | --- | --- | +| Router | BF16 `x`, `W` upcast to FP32; one pinned vLLM batch-invariant GEMM (IEEE FP32) | none | +| Softmax | max subtraction, `exp`, row sum by a fixed pairwise tree (`fixed_order_row_sum`) | none | +| Selection | `argsort(logits, descending, stable)[:, :10]`: score descending, equal scores to the lower expert id | none | +| Weights | the ten selected FP32 probabilities divided by their fixed-tree sum | none | +| Experts | per (token, route) row: `gate_up` GEMM (BF16 in, FP32 accumulate, BF16 out), SwiGLU in FP32, BF16 cast, `down` GEMM | none: one writer per row | +| Combine | `out = e0*w0`, then `out += e_k*w_k` for k = 1..9 in route order, FP32; one BF16 cast | none | +| Shared expert (TP4 block) | `routed + shared * sigmoid(gate)` in BF16, then one TP all-reduce | collective as configured | + +Payload dtype between stages: BF16 expert outputs, FP32 route weights, BF16 +block output. The expert GEMMs are vLLM 0.30.0's `fused_moe` Triton kernel with +the tile vLLM itself uses under `VLLM_BATCH_INVARIANT=1` +(`BLOCK_SIZE_M/N/K = 64/64/32`, `GROUP_SIZE_M = 8`, `SPLIT_K = 1`), passed +explicitly so no environment variable, tuned-config file or token count can +change it. The routed weight is *not* multiplied inside the kernel and the +kernel does not sum over routes; that stays in the combine above. + +Fail-closed checks: CUDA tensors only (no CPU fallback, no ROCm), BF16 weights of +the exact `[512, 2*I, H]` / `[512, H, I]` shapes, finite FP32 router logits, +vLLM exactly 0.30.0, `VLLM_TRITON_USE_TD` off (the tensor-descriptor path is a +different instruction stream), and IEEE Triton FP32. + +### Backward + +`_RoutedExperts.backward` recomputes the `[gate, up]` projection with the same +grouped kernel, then visits experts in ascending id. Each expert's `dgate_up` +and `ddown` slice is one pinned GEMM written once into a dense buffer; `dx` is +accumulated per token in ascending expert order (a token's ten experts are +distinct, so each launch has unique rows). Router-weight gradients flow through +the selected probabilities; the indices are discrete. + +## Reuse decision + +| Implementation | Routes BI | Output BI | Backward | Why not reused as is | +| --- | --- | --- | --- | --- | +| vLLM `fused_topk` + `fused_experts`, `VLLM_BATCH_INVARIANT=0` | no | no | none | tuned configs and the BF16 router GEMM depend on the token count | +| vLLM, `VLLM_BATCH_INVARIANT=1` | yes | yes | none | no VJP for training; BF16 router logits (Qwen3-Next and VIME route in FP32); routed weight applied in the expert GEMM epilogue on the BF16 output, then routes summed by `moe_sum` | +| HF `Qwen3NextExperts` (eager) | yes | yes | autograd | `dx`/`dW` change with batch size (cuBLAS shape heuristics); BF16 router; `index_add_` combine in BF16 | +| FlashInfer `cutlass_fused_moe` | (torch routing) | no | none | output changes with batch size; no VJP | +| Megatron-core 0.16 `MoELayer` + TE 2.16 grouped GEMM, as VIME configures Qwen3-Next | no (256+) | no (1024) | autograd, not BI | FP32 router routes like FP64, but the routing, output, dx and dW change with batch size | +| SGLang 0.5.21 Triton `fused_moe`, default | no (256+) | no (256+) | none | tuned configs and the BF16 router GEMM depend on the token count | +| SGLang, deterministic inference | yes | yes | none | inference only; BF16 router logits; its deterministic tile (64/64/32) is the one vLLM uses in batch-invariant mode | + +What RL-Kernel reuses: vLLM's batch-invariant GEMM for the router and the +backward, and vLLM's `fused_moe` Triton kernel for the forward expert GEMMs, +which is batch-invariant per route at the fixed tile. What it implements: the +FP32 fixed-order routing, the stable tie rule, the atomic-free FP32 combine, the +deterministic backward and the TP4 boundaries. The grouped kernel's per-route +output is bitwise equal to the per-expert pinned GEMM loop it replaced +(`tests/models/qwen3_next/check_qwen3_next_forward.py::test_grouped_*`), so this is a speed change +with no change in bits. + +## Results + +Measured on B200 with `tools/validation/models/qwen3_next_moe_prior_art.py` and the TP4 gate +`tools/validation/models/qwen3_next_tp_moe_check.py`; the reports, the figure and the exact +commands are in +[`docs/usage/evidence/qwen3-next-moe-route-b200/`](../usage/evidence/qwen3-next-moe-route-b200/README.md). + +![Qwen3-Next routed MoE prior art](../usage/evidence/qwen3-next-moe-route-b200/figure.png) + +TP1 shape (H=2048, 512 experts, top-10, width 512), random weights, BF16: + +| | routes / output BI (16-1024 tokens) | dx / dW BI | rel. L2 vs FP64 | tokens routed unlike FP64 (of 256) | forward, 1 / 64 / 1024 / 4096 tokens (ms) | fwd + bwd, 64 / 1024 (ms) | +| --- | --- | --- | --- | --- | --- | --- | +| **RL-Kernel `shared_moe`** | **yes / yes** | **yes / yes** | **3.9e-3** | **0** | 1.39 / 1.69 / 2.04 / 3.05 | 114 / 169 | +| Megatron-core + TE (VIME config) | no / no | no / no (1024) | 4.6e-3 | 0 | 3.8 / 9.7 / 11.7 / 12.1 | 45 / 51 | +| vLLM BI=1 | yes / yes | no backward | 7.0e-2 | 11 | 0.39 / 0.66 / 0.87 / 1.05 | - | +| SGLang deterministic | yes / yes | no backward | 7.0e-2 | 11 | 0.35 / 0.66 / 0.84 / 1.37 | - | +| vLLM BI=0 | no / no | no backward | 7.0e-2 | 11 | 0.38 / 0.66 / 0.87 / 1.06 | - | +| SGLang default | no / no | no backward | 7.0e-2 | 11 | 0.34 / 0.61 / 0.83 / 1.35 | - | +| FlashInfer CUTLASS | no / no | no backward | 7.0e-2 | 11 | 0.37 / 0.72 / 0.90 / 1.12 | - | +| HF eager | yes / yes | no / no (1024) | 7.0e-2 | 11 | 2.1 / 62 / 89 / 87 | 877 / 1284 | + +* The accuracy gap is the router: every candidate with BF16 router logits sends + 11 of 256 tokens to a different expert set than FP64 routing does. RL-Kernel + and Megatron-core (VIME's training path) route in FP32 and select the same + experts as FP64. +* The only batch-invariant candidates are RL-Kernel and the inference-only + modes of vLLM and SGLang; of these, only RL-Kernel has a backward, so only it + can run the same forward on the training and the rollout side. +* RL-Kernel's forward is 2-4x slower than vLLM's. About 0.5 ms of it, at any + token count, is the FP32 IEEE router GEMM, which is part of the contract. + The expert GEMMs themselves run in vLLM's kernel. The per-expert GEMM loop + they replace (same bits) took 47 / 66 / 71 ms at 64 / 1024 / 4096 tokens in + a local B200 run of the same runner; that run is not part of this evidence. +* The backward is still a per-expert loop of pinned GEMMs (512 experts x 4 + GEMMs): 169 ms at 1024 tokens, against 51 ms for Megatron-core + TE, which is + not batch-invariant. It dominates a training step's MoE time. + +TP4 gate on the real layer-0 weights (`tp4-moe/rank-*.json`): all twelve cases +pass on all four ranks, including the round trip of all 512 experts, identical +routes on every rank, full == chunked and reordered batches at 8-1024 tokens, +and identical replicated router gradients. The TP4 block forward (routed + +shared expert + all-reduce) takes 1.9-2.4 ms at 8-1024 tokens. + +## Tests + +| File | Device | Covers | +| --- | --- | --- | +| `tests/models/qwen3_next/test_qwen3_next_forward_contract.py` | CPU | CPU rejection, vLLM pin, fixed-order sum/softmax row independence | +| `tests/models/qwen3_next/test_qwen3_next_tp_blocks.py` | CPU | HF shard/assemble round trips, replica drift rejection, TP4 parameter ownership | +| `tests/validation/common/test_tensor_identity.py` | CPU | raw-bit identity (signed zero, NaN/Inf, dtype) | +| `tests/models/qwen3_next/check_qwen3_next_forward.py` | CUDA + vLLM | GEMM batch/chunk/reorder, route ties, combine order, MoE VJP vs FP64, route and output batch invariance, grouped == per-expert bitwise, fail-closed TD path | +| `tools/validation/models/qwen3_next_tp_moe_check.py` | 4 x CUDA | real layer-0 weights: HF round trip, replicated routes, chunk/reorder, training forward, gradients | + +The `check_` file imports vLLM, so it runs in the `Qwen3-Next-provider-GPU` +workflow rather than the default collection. + +## Scope + +TP4 with EP1 (all 512 experts on every rank) on CUDA B200. ROCm, EP > 1 and +other TP sizes are separate claims. Model-level parity uses this block but is +the `full_model_chain` row. 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-moe-route-b200/README.md b/docs/usage/evidence/qwen3-next-moe-route-b200/README.md new file mode 100644 index 000000000..0a86096b4 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/README.md @@ -0,0 +1,47 @@ +# Qwen3-Next routed MoE: prior art and TP4 gate (B200) + +`report.json` and `figure.png` were measured at commit `c58816b` from a clean +clone (`tracked_tree_dirty: false`) on one B200 (driver 580.126.20, torch +2.13.0+cu130, Triton 3.7.1), in two Python environments because no single one +can import every candidate: + +* the training environment (VIME 0.3.2 with Megatron-core 0.16.0rc0 and + Transformer Engine 2.16.1, vLLM 0.30.0, FlashInfer 0.6.18.post1, + transformers 5.17.0): RL-Kernel, HF, vLLM, FlashInfer, Megatron-core + TE; +* a second environment with SGLang 0.5.21: SGLang in default and + deterministic-inference mode. Its compiled `sgl_kernel` 0.3.21 is built for + another libtorch ABI and cannot load next to torch 2.13, so SGLang's two + kernels on this path are replaced by SGLang's own implementations of the same + operations (Triton `moe_sum_reduce`, JIT `moe_align_block_size`); calling any + other `sgl_kernel` symbol fails. + +`environments` in `report.json` records both. The TP4 gate files were measured +at `2e950f3`, whose MoE code is identical (`c58816b` only adds runner candidates). + +| File | Produced by | +| --- | --- | +| `report.json` | `TRITON_F32_DEFAULT=ieee python tools/validation/models/qwen3_next_moe_prior_art.py --only rl_kernel_cuda,hf_transformers,vllm_bi0,vllm_bi1,flashinfer_cutlass,megatron_te --out main.json` (training environment), then `... --only sglang_triton,sglang_deterministic --merge main.json --out report.json` (SGLang environment); one GPU | +| `figure.png` | `python tools/validation/models/plot_qwen3_next_moe_prior_art.py report.json` | +| `tp4-moe/rank-{0..3}.json` | `TRITON_F32_DEFAULT=ieee torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_moe_check.py --checkpoint --output tp4-moe` | + +The checkpoint is the official `Qwen/Qwen3-Next-80B-A3B-Instruct` revision +`9c7f2fbe84465e40164a94cc16cd30b6999b0cc7`; the TP4 gate reads layer 0's MoE. + +In the same job, `TRITON_F32_DEFAULT=ieee python -m pytest -q +tests/models/qwen3_next/check_qwen3_next_forward.py tests/validation/common/test_tensor_identity.py +tests/models/qwen3_next/test_qwen3_next_forward_contract.py tests/models/qwen3_next/test_qwen3_next_tp_blocks.py` +passed (65 tests). `tests/integrations/common/test_framework_operator_integrations.py` was collected +into that same process and failed, as it must once a `check_` file has imported +vLLM; run in its own process, as CI does, it passes. + +## Reading the report + +* Batch invariance: 8 probe tokens computed alone, and placed first and last in + batches of 16, 64, 256 and 1024 tokens; `true` means bitwise equal everywhere. + `dweight_zero_rows_bitwise` appends rows whose output gradient is zero. +* Accuracy: 256 tokens against the HF formula evaluated in FP64 with FP64 + routing; `tokens_with_different_expert_set` counts tokens whose ten selected + experts differ from the FP64 selection. +* Latency: CUDA-event medians, candidates interleaved with the order reversed + every iteration. Weights are random (scale 0.02), shape H=2048, 512 experts, + top-10, expert width 512 (TP1). diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/figure.png b/docs/usage/evidence/qwen3-next-moe-route-b200/figure.png new file mode 100644 index 000000000..e09ec93ac Binary files /dev/null and b/docs/usage/evidence/qwen3-next-moe-route-b200/figure.png differ diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/report.json b/docs/usage/evidence/qwen3-next-moe-route-b200/report.json new file mode 100644 index 000000000..5a9fa5056 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/report.json @@ -0,0 +1,381 @@ +{ + "kind": "qwen3_next_prior_art_report", + "rfc": 428, + "rl_kernel_commit": "c58816bf54940195b543beda14e0de65a71182c6", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": null, + "sgl-kernel": null, + "megatron_core": "0.16.0rc0", + "megatron-core": "0.16.0rc0", + "transformer_engine": "2.16.1+c9877beb", + "vime": "0.3.2" + } + }, + "op": "qwen3_next_routed_moe", + "shape": { + "hidden": 2048, + "experts": 512, + "top_k": 10, + "expert_width": 512 + }, + "probe_rows": 8, + "candidates": [ + { + "key": "rl_kernel_cuda", + "name": "rl-kernel shared_moe", + "source": "rl_engine.integrations.qwen3_next_forward.shared_moe", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "repeatable": true, + "dx_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "dweight_zero_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "accuracy": { + "rel_l2_vs_fp64": 0.0038748319159183277, + "max_abs_vs_fp64": 0.0021274995396960983, + "tokens_with_different_expert_set": 0, + "tokens": 256 + }, + "batch_invariant": true + }, + { + "key": "hf_transformers", + "name": "HF transformers Qwen3NextExperts (eager)", + "source": "transformers 5.17.0 modeling_qwen3_next", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "repeatable": true, + "dx_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": false + }, + "dweight_zero_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": false + }, + "accuracy": { + "rel_l2_vs_fp64": 0.0700071926559018, + "max_abs_vs_fp64": 0.10691143625663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": false + }, + { + "key": "vllm_bi0", + "name": "vLLM fused_topk + fused_experts (Triton), VLLM_BATCH_INVARIANT=0", + "source": "vllm 0.30.0 model_executor/layers/fused_moe", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.06987086350006304, + "max_abs_vs_fp64": 0.10691143625663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": false + }, + { + "key": "vllm_bi1", + "name": "vLLM fused_topk + fused_experts (Triton), VLLM_BATCH_INVARIANT=1", + "source": "vllm 0.30.0 model_executor/layers/fused_moe", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.06987086350006304, + "max_abs_vs_fp64": 0.10691143625663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": true + }, + { + "key": "flashinfer_cutlass", + "name": "FlashInfer cutlass_fused_moe (routes from torch.topk)", + "source": "flashinfer 0.6.18.post1 fused_moe.cutlass_fused_moe", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.06984271064260991, + "max_abs_vs_fp64": 0.10739971750663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": false + }, + { + "key": "megatron_te", + "name": "Megatron-core MoELayer + TE grouped GEMM (VIME's Qwen3-Next config)", + "source": "megatron-core 0.16.0rc0, transformer-engine 2.16.1+c9877beb", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": false + }, + "repeatable": true, + "dx_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": false + }, + "dweight_zero_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": false + }, + "accuracy": { + "rel_l2_vs_fp64": 0.00456482750054121, + "max_abs_vs_fp64": 0.0022296928321409726, + "tokens_with_different_expert_set": 0, + "tokens": 256 + }, + "batch_invariant": false + }, + { + "key": "sglang_triton", + "name": "SGLang fused_moe (Triton), default", + "source": "sglang 0.5.21 moe_runner/triton_utils", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": false, + "1024": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.06984328047814348, + "max_abs_vs_fp64": 0.10739971750663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": false + }, + { + "key": "sglang_deterministic", + "name": "SGLang fused_moe (Triton), deterministic inference", + "source": "sglang 0.5.21 moe_runner/triton_utils", + "route_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "output_rows_bitwise": { + "16": true, + "64": true, + "256": true, + "1024": true + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.06984328047814348, + "max_abs_vs_fp64": 0.10739971750663532, + "tokens_with_different_expert_set": 11, + "tokens": 256 + }, + "batch_invariant": true + } + ], + "latency": { + "forward_us": { + "1": { + "rl_kernel_cuda": 1386.0480189323425, + "hf_transformers": 2066.7200088500977, + "vllm_bi0": 374.9440014362335, + "vllm_bi1": 387.91999220848083, + "flashinfer_cutlass": 368.51200461387634, + "megatron_te": 3787.440061569214, + "sglang_triton": 339.1679972410202, + "sglang_deterministic": 349.7439920902252 + }, + "8": { + "rl_kernel_cuda": 1395.039975643158, + "hf_transformers": 13072.400093078613, + "vllm_bi0": 381.9680064916611, + "vllm_bi1": 385.328009724617, + "flashinfer_cutlass": 432.2880059480667, + "megatron_te": 5061.583995819092, + "sglang_triton": 332.8000009059906, + "sglang_deterministic": 353.5839915275574 + }, + "64": { + "rl_kernel_cuda": 1687.5839829444885, + "hf_transformers": 62156.3835144043, + "vllm_bi0": 660.0319743156433, + "vllm_bi1": 658.3200097084045, + "flashinfer_cutlass": 723.4559953212738, + "megatron_te": 9674.047946929932, + "sglang_triton": 606.8799793720245, + "sglang_deterministic": 660.75199842453 + }, + "256": { + "rl_kernel_cuda": 1846.1920022964478, + "hf_transformers": 85269.61517333984, + "vllm_bi0": 779.8559963703156, + "vllm_bi1": 790.6079888343811, + "flashinfer_cutlass": 858.7839901447296, + "megatron_te": 11665.247917175293, + "sglang_triton": 759.0239942073822, + "sglang_deterministic": 819.1200196743011 + }, + "1024": { + "rl_kernel_cuda": 2043.7918901443481, + "hf_transformers": 89267.74215698242, + "vllm_bi0": 874.2719888687134, + "vllm_bi1": 868.9280152320862, + "flashinfer_cutlass": 902.5439918041229, + "megatron_te": 11738.048076629639, + "sglang_triton": 833.1039845943451, + "sglang_deterministic": 841.6959941387177 + }, + "4096": { + "rl_kernel_cuda": 3053.2480478286743, + "hf_transformers": 87444.65637207031, + "vllm_bi0": 1059.056043624878, + "vllm_bi1": 1051.2160062789917, + "flashinfer_cutlass": 1124.2560148239136, + "megatron_te": 12122.560024261475, + "sglang_triton": 1346.0160493850708, + "sglang_deterministic": 1367.3120141029358 + } + }, + "forward_plus_backward_us": { + "64": { + "rl_kernel_cuda": 113590.37017822266, + "hf_transformers": 876752.3193359375, + "megatron_te": 44856.64176940918 + }, + "1024": { + "rl_kernel_cuda": 169058.31909179688, + "hf_transformers": 1284274.658203125, + "megatron_te": 50968.2559967041 + } + } + }, + "environments": { + "main": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": null, + "sgl-kernel": null, + "megatron_core": "0.16.0rc0", + "megatron-core": "0.16.0rc0", + "transformer_engine": "2.16.1+c9877beb", + "vime": "0.3.2" + } + }, + "extra": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": "0.5.21", + "sgl-kernel": "0.3.21", + "megatron_core": null, + "megatron-core": null, + "transformer_engine": "2.20.2", + "vime": null + } + } + } +} diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-0.json b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-0.json new file mode 100644 index 000000000..188c00b77 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-0.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 0, + "rl_kernel_commit": "2e950f386bd3c4c2c07f41472f48878be94ea3b8", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 4.598405212163925e-09, + "experts.down": 7.8580342233181e-09, + "shared_expert.gate_proj.weight": 2.5890767574310303e-07, + "shared_expert.up_proj.weight": 4.3585896492004395e-07, + "shared_expert.down_proj.weight": 2.9243528842926025e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2185.663938522339, + "64": 2399.2319107055664, + "256": 1903.6799669265747, + "1024": 1860.6079816818237 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-1.json b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-1.json new file mode 100644 index 000000000..755bed167 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-1.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 1, + "rl_kernel_commit": "2e950f386bd3c4c2c07f41472f48878be94ea3b8", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.240994945168495e-09, + "experts.down": 5.558831617236137e-09, + "shared_expert.gate_proj.weight": 3.129243850708008e-07, + "shared_expert.up_proj.weight": 3.725290298461914e-07, + "shared_expert.down_proj.weight": 4.6938657760620117e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2204.67209815979, + "64": 2400.2881050109863, + "256": 1920.4479455947876, + "1024": 1854.464054107666 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-2.json b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-2.json new file mode 100644 index 000000000..72343dc29 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-2.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 2, + "rl_kernel_commit": "2e950f386bd3c4c2c07f41472f48878be94ea3b8", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.0227162167429924e-09, + "experts.down": 5.122274160385132e-09, + "shared_expert.gate_proj.weight": 2.337619662284851e-07, + "shared_expert.up_proj.weight": 3.9674341678619385e-07, + "shared_expert.down_proj.weight": 2.384185791015625e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2204.67209815979, + "64": 2398.303985595703, + "256": 1881.0559511184692, + "1024": 1856.511950492859 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-3.json b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-3.json new file mode 100644 index 000000000..2c4d8f376 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-moe-route-b200/tp4-moe/rank-3.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 3, + "rl_kernel_commit": "2e950f386bd3c4c2c07f41472f48878be94ea3b8", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.546585164964199e-09, + "experts.down": 5.296897143125534e-09, + "shared_expert.gate_proj.weight": 1.2665987014770508e-07, + "shared_expert.up_proj.weight": 2.644956111907959e-07, + "shared_expert.down_proj.weight": 2.998858690261841e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2206.4640522003174, + "64": 2399.104118347168, + "256": 1923.0719804763794, + "1024": 1855.455994606018 + } + } +} 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/docs/usage/evidence/qwen3-next-tp4-mixers-b200/README.md b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/README.md new file mode 100644 index 000000000..a21252c17 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/README.md @@ -0,0 +1,71 @@ +# Qwen3-Next TP4 mixers: real-weight gates and attention prior art (B200) + +| Directory | Measured at | Produced by | +| --- | --- | --- | +| `tp4-gdn/` | `9e6c4f2` | `torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_gdn_check.py --checkpoint --output tp4-gdn` | +| `tp4-attention/` | `9e6c4f2` | `torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_attention_check.py --checkpoint --output tp4-attention` | +| `tp4-moe/` | `9e6c4f2` | `TRITON_F32_DEFAULT=ieee torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_moe_check.py --checkpoint --output tp4-moe` | +| `attention/report.json`, `attention/figure.png` | `703c6dc` | two processes, then `tools/validation/models/plot_qwen3_next_attention_prior_art.py attention/report.json` (see below) | + +Every file records its commit and `tracked_tree_dirty: false`; each run used a +clean clone with its own `rl_engine._C` built in place, on B200 GPUs (driver +580.126.20), torch 2.13.0+cu130, vLLM 0.30.0, NCCL 2.29.7. `` is the +official `Qwen/Qwen3-Next-80B-A3B-Instruct` revision +`9c7f2fbe84465e40164a94cc16cd30b6999b0cc7`: the GDN and MoE gates read layer 0, +the attention gate layer 3. `703c6dc` adds only the attention prior-art runner +and plot on top of `9e6c4f2` (and the MoE evidence merge). + +In the job that ran the gates, the mixer GPU checks +(`tests/check_qwen3_next_{attention,gdn_bridge,gdn_sequence,conv_bridge,core_matrix,shared_core,forward}.py`) +and the CPU tests passed together (123 tests); `tests/integrations/common/test_framework_operator_integrations.py`, +collected into the same process after vLLM was imported, failed as it must and +passes in its own process. The existing attention suites +(`test_deterministic_attention_cuda`, `test_attention`, `test_attention_correctness`, +`test_kv_cache_attention`, `test_attention_contract`, `test_attention_dispatch`) +gave 766 passed and the two `test_attention_dispatch` failures that `main` also has. + +## Attention prior-art report + +One TP4 rank of Qwen3-Next full attention: 4 query heads, 1 KV head, D=256, +scale 1/16, causal. A 1000-token target sequence is computed alone and in a +batch with sequences of 777, 1500 and 64 tokens (first and last); +`prefill_decode_bitwise` compares the last 64 rows and the last row computed +against the full KV with the same rows of the full prefill; +`backward_batch_bitwise` compares the target's dq/dk/dv computed alone with the +same sequence packed first and last among the companions (one launch where the +engine has a packed backward). Accuracy is against FP64; latencies are CUDA-event medians with +the candidate order reversed every iteration. + +The report was produced in two processes of the same job and merged: + +```bash +python tools/validation/models/qwen3_next_attention_prior_art.py \ + --only rl_kernel_cuda,torch_sdpa,vllm_fa2_auto,vllm_fa2_split1,vllm_triton_2d,flashinfer,fa4_cute,te_training,megatron_local \ + --out report-main.json +python tools/validation/models/qwen3_next_attention_prior_art.py --only te_fused --merge report-main.json --out report.json +``` + +Both run in the training environment (vLLM 0.30.0, Megatron-core 0.16.0rc0, +Transformer Engine 2.16.1, cuDNN 9.20) with `CUDNN_HOME` pointing at the +`nvidia-cudnn-cu13` wheel. Notes on the engines added in this revision: + +- **FlashAttention-4** is the CuTe-DSL build vendored by vLLM + (`vllm.vllm_flash_attn.cute`), varlen forward and its default backward. Its + `deterministic=True` backward asserts for head dim 256 on SM100. +- **Transformer Engine**: with cuDNN 9.20 on SM100, TE offers its fused cuDNN + backend for head dim 256 only in inference mode; in training mode it selects + the unfused PyTorch fallback (`NVTE_DEBUG=1` shows the selection). Both are + measured. The unfused fallback raises a shape error when the query is shorter + than the KV (`padding_causal_bottom_right`, THD), recorded as + `prefill_decode_failed`. `NVTE_ALLOW_NONDETERMINISTIC_ALGO` only affects the + fused backward, which does not exist here, so it is not a separate row. +- The fused TE row runs in its own process because cudnn-frontend's runtime + loader refuses to start when both `libcudart.so.12` and `libcudart.so.13` + can be `dlopen`ed, and the compute nodes' system image provides a CUDA 12 + runtime next to the CUDA 13 one this stack uses. That process had a + non-loadable `libcudart.so.12` stub first on `LD_LIBRARY_PATH`; only + `libcudart.so.13` was loaded in either process. Its latencies come from that + process, not interleaved with the other candidates. +- **Megatron-core**'s own `DotProductAttention` (not the TE extension) has no + packed-sequence path, so it runs one launch per sequence with a bottom-right + causal mask; its batch invariance follows from that. diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/figure.png b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/figure.png new file mode 100644 index 000000000..cb0c65b02 Binary files /dev/null and b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/figure.png differ diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/report.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/report.json new file mode 100644 index 000000000..3d55ea2aa --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/report.json @@ -0,0 +1,346 @@ +{ + "kind": "qwen3_next_prior_art_report", + "rfc": 428, + "rl_kernel_commit": "703c6dca23da2f42b288fefe70f2073ef6a97986", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": null, + "megatron-core": "0.16.0rc0", + "megatron_core": "0.16.0rc0", + "transformer_engine": "2.16.1+c9877beb", + "vime": "0.3.2" + } + }, + "op": "qwen3_next_tp4_attention", + "shape": { + "q_heads": 4, + "kv_heads": 1, + "head_dim": 256, + "scale": 0.0625 + }, + "target_tokens": 1000, + "companion_tokens": [ + 777, + 1500, + 64 + ], + "candidates": [ + { + "key": "rl_kernel_cuda", + "name": "RL-Kernel deterministic attention (one launch per sequence)", + "source": "rl_engine.integrations.qwen3_next_forward.shared_attention", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": true + }, + "repeatable": true, + "backward_batch_bitwise": { + "dq": true, + "dk": true, + "dv": true + }, + "accuracy": { + "rel_l2_vs_fp64": 0.0016090058011337504, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": true + }, + { + "key": "torch_sdpa", + "name": "torch SDPA (default backend, one launch per sequence)", + "source": "torch 2.13.0+cu130 scaled_dot_product_attention", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": false, + "last_1": false + }, + "repeatable": true, + "backward_batch_bitwise": { + "dq": false, + "dk": true, + "dv": true + }, + "accuracy": { + "rel_l2_vs_fp64": 0.0019992834566479067, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": false + }, + { + "key": "vllm_fa2_auto", + "name": "vLLM FlashAttention-2 varlen, num_splits=0", + "source": "vllm 0.30.0 vllm_flash_attn.flash_attn_varlen_func", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.0020096983122931318, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": false + }, + { + "key": "vllm_fa2_split1", + "name": "vLLM FlashAttention-2 varlen, num_splits=1", + "source": "vllm 0.30.0 vllm_flash_attn.flash_attn_varlen_func", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": true + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.0020096983122931318, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": true + }, + { + "key": "vllm_triton_2d", + "name": "vLLM Triton unified attention (2D kernel)", + "source": "vllm 0.30.0 v1/attention/ops/triton_unified_attention", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.002021002108737795, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": false + }, + { + "key": "flashinfer", + "name": "FlashInfer BatchPrefillWithRaggedKVCache", + "source": "flashinfer 0.6.18.post1", + "batch_bitwise": { + "first": false, + "last": false + }, + "prefill_decode_bitwise": { + "last_64": false, + "last_1": false + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.002146056864137061, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": false + }, + { + "key": "fa4_cute", + "name": "FlashAttention-4 (CuTe DSL) varlen, vendored by vLLM", + "source": "vllm 0.30.0 vllm_flash_attn/cute (its head-dim-256 SM100 backward has no deterministic mode)", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": true + }, + "repeatable": true, + "backward_batch_bitwise": { + "dq": true, + "dk": true, + "dv": true + }, + "accuracy": { + "rel_l2_vs_fp64": 0.0020224522459116262, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": true + }, + { + "key": "te_training", + "name": "Transformer Engine DotProductAttention (THD), training: unfused fallback", + "source": "transformer-engine 2.16.1+c9877beb (selected for head-dim-256 training on SM100)", + "batch_bitwise": { + "first": false, + "last": false + }, + "prefill_decode_failed": "RuntimeError: The size of tensor a (1000) must match the size of tensor b (64) at non-singleton dimension 3", + "prefill_decode_bitwise": { + "last_64": false, + "last_1": false + }, + "repeatable": true, + "backward_batch_bitwise": { + "dq": false, + "dk": false, + "dv": false + }, + "accuracy": { + "rel_l2_vs_fp64": 0.003577689514635195, + "max_abs_vs_fp64": 0.012889566535712937 + }, + "batch_invariant": false + }, + { + "key": "megatron_local", + "name": "Megatron-core DotProductAttention (local, one launch per sequence)", + "source": "megatron-core 0.16.0rc0", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": true + }, + "repeatable": true, + "backward_batch_bitwise": { + "dq": true, + "dk": true, + "dv": true + }, + "accuracy": { + "rel_l2_vs_fp64": 0.003577688032086663, + "max_abs_vs_fp64": 0.012889566535712937 + }, + "batch_invariant": true + }, + { + "key": "te_fused", + "name": "Transformer Engine DotProductAttention (THD), inference: cuDNN fused", + "source": "transformer-engine 2.16.1+c9877beb (no fused head-dim-256 backward on SM100)", + "batch_bitwise": { + "first": true, + "last": true + }, + "prefill_decode_bitwise": { + "last_64": true, + "last_1": true + }, + "repeatable": true, + "accuracy": { + "rel_l2_vs_fp64": 0.002051519013469327, + "max_abs_vs_fp64": 0.0066821561991154965 + }, + "batch_invariant": true + } + ], + "latency": { + "prefill_us": { + "512": { + "rl_kernel_cuda": 725.02401471138, + "torch_sdpa": 66.0800002515316, + "vllm_fa2_auto": 118.46399679780006, + "vllm_fa2_split1": 118.71999874711037, + "vllm_triton_2d": 263.88800144195557, + "flashinfer": 237.88800090551376, + "fa4_cute": 210.09600162506104, + "te_training": 2039.6959781646729, + "megatron_local": 324.5760053396225, + "te_fused": 336.70400083065033 + }, + "2048": { + "rl_kernel_cuda": 10141.616344451904, + "torch_sdpa": 115.9840002655983, + "vllm_fa2_auto": 152.8479978442192, + "vllm_fa2_split1": 156.78400546312332, + "vllm_triton_2d": 324.6240019798279, + "flashinfer": 293.16800832748413, + "fa4_cute": 211.31199598312378, + "te_training": 5458.016157150269, + "megatron_local": 320.8480030298233, + "te_fused": 323.0559974908829 + }, + "8192": { + "rl_kernel_cuda": 159770.19500732422, + "torch_sdpa": 607.344001531601, + "vllm_fa2_auto": 661.2959802150726, + "vllm_fa2_split1": 657.3439836502075, + "vllm_triton_2d": 1134.160041809082, + "flashinfer": 764.3359899520874, + "fa4_cute": 312.0640069246292, + "te_training": 22656.047821044922, + "megatron_local": 3245.0079917907715, + "te_fused": 612.496018409729 + } + }, + "decode_us_kv8192": { + "rl_kernel_cuda": 757.8720152378082, + "torch_sdpa": 314.2559975385666, + "vllm_fa2_auto": 164.39999639987946, + "vllm_fa2_split1": 336.4799916744232, + "vllm_triton_2d": 564.3199980258942, + "flashinfer": 235.24799942970276, + "fa4_cute": 256.49599730968475, + "te_training": 4540.832042694092, + "megatron_local": 304.56000566482544, + "te_fused": 443.83999705314636 + }, + "forward_plus_backward_us_2048": { + "rl_kernel_cuda": 25026.527404785156, + "torch_sdpa": 741.6639924049377, + "fa4_cute": 912.8959774971008, + "te_training": 10321.47216796875, + "megatron_local": 1336.4800214767456 + } + }, + "environments": { + "main": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": null, + "megatron-core": "0.16.0rc0", + "megatron_core": "0.16.0rc0", + "transformer_engine": "2.16.1+c9877beb", + "vime": "0.3.2" + } + }, + "extra": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "cuda": "13.0", + "libraries": { + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "transformers": "5.17.0", + "triton": "3.7.1", + "sglang": null, + "megatron-core": "0.16.0rc0", + "megatron_core": "0.16.0rc0", + "transformer_engine": "2.16.1+c9877beb", + "vime": "0.3.2" + } + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-0.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-0.json new file mode 100644 index 000000000..457309564 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-0.json @@ -0,0 +1,36 @@ +{ + "status": "passed", + "scope": "tp4_full_attention_layer3", + "rank": 0, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "attention": { + "cases": [ + "hf-load-export", + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "batch4-varlen-reorder", + "sixteen-token-decode", + "prompt-gradient", + "kv-pair-gradient", + "qk-norm-gradient" + ], + "gradient_max_abs": { + "q_proj.weight": 7.748603820800781e-06, + "k_proj.weight": 1.52587890625e-05, + "v_proj.weight": 3.075599670410156e-05, + "o_proj.weight": 1.2099742889404297e-05, + "q_norm.weight": 1.8835067749023438e-05, + "k_norm.weight": 1.8715858459472656e-05 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-1.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-1.json new file mode 100644 index 000000000..c8b9c07a5 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-1.json @@ -0,0 +1,36 @@ +{ + "status": "passed", + "scope": "tp4_full_attention_layer3", + "rank": 1, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "attention": { + "cases": [ + "hf-load-export", + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "batch4-varlen-reorder", + "sixteen-token-decode", + "prompt-gradient", + "kv-pair-gradient", + "qk-norm-gradient" + ], + "gradient_max_abs": { + "q_proj.weight": 8.761882781982422e-06, + "k_proj.weight": 1.52587890625e-05, + "v_proj.weight": 3.075599670410156e-05, + "o_proj.weight": 1.2099742889404297e-05, + "q_norm.weight": 1.8835067749023438e-05, + "k_norm.weight": 1.8715858459472656e-05 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-2.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-2.json new file mode 100644 index 000000000..db5ec61cb --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-2.json @@ -0,0 +1,36 @@ +{ + "status": "passed", + "scope": "tp4_full_attention_layer3", + "rank": 2, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "attention": { + "cases": [ + "hf-load-export", + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "batch4-varlen-reorder", + "sixteen-token-decode", + "prompt-gradient", + "kv-pair-gradient", + "qk-norm-gradient" + ], + "gradient_max_abs": { + "q_proj.weight": 1.633167266845703e-05, + "k_proj.weight": 2.1338462829589844e-05, + "v_proj.weight": 2.3126602172851562e-05, + "o_proj.weight": 2.2292137145996094e-05, + "q_norm.weight": 1.8835067749023438e-05, + "k_norm.weight": 1.8715858459472656e-05 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-3.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-3.json new file mode 100644 index 000000000..5ee714ab1 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-attention/rank-3.json @@ -0,0 +1,36 @@ +{ + "status": "passed", + "scope": "tp4_full_attention_layer3", + "rank": 3, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "attention": { + "cases": [ + "hf-load-export", + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "batch4-varlen-reorder", + "sixteen-token-decode", + "prompt-gradient", + "kv-pair-gradient", + "qk-norm-gradient" + ], + "gradient_max_abs": { + "q_proj.weight": 1.2993812561035156e-05, + "k_proj.weight": 2.1338462829589844e-05, + "v_proj.weight": 2.3126602172851562e-05, + "o_proj.weight": 9.894371032714844e-06, + "q_norm.weight": 1.8835067749023438e-05, + "k_norm.weight": 1.8715858459472656e-05 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-0.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-0.json new file mode 100644 index 000000000..2e2501dfb --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-0.json @@ -0,0 +1,32 @@ +{ + "status": "passed", + "scope": "tp4_gdn_layer0", + "rank": 0, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "cases": [ + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "prompt-gradient", + "replicated-norm-gradient", + "hf-load-export" + ], + "gradient_max_abs": { + "A_log": 4.462208380573429e-12, + "dt_bias": 3.595346242946107e-12, + "in_proj_qkvz.weight": 6.948539521545172e-10, + "in_proj_ba.weight": 3.6834535421803594e-11, + "conv1d.weight": 3.128661774098873e-09, + "norm.weight": 3.91155481338501e-08, + "out_proj.weight": 3.255991032347083e-10 + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-1.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-1.json new file mode 100644 index 000000000..33bda4b32 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-1.json @@ -0,0 +1,32 @@ +{ + "status": "passed", + "scope": "tp4_gdn_layer0", + "rank": 1, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "cases": [ + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "prompt-gradient", + "replicated-norm-gradient", + "hf-load-export" + ], + "gradient_max_abs": { + "A_log": 1.0550138540565968e-10, + "dt_bias": 9.458744898438454e-11, + "in_proj_qkvz.weight": 3.055902197957039e-10, + "in_proj_ba.weight": 3.296918293926865e-11, + "conv1d.weight": 9.19681042432785e-09, + "norm.weight": 3.91155481338501e-08, + "out_proj.weight": 1.8553691916167736e-10 + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-2.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-2.json new file mode 100644 index 000000000..c3d11374b --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-2.json @@ -0,0 +1,32 @@ +{ + "status": "passed", + "scope": "tp4_gdn_layer0", + "rank": 2, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "cases": [ + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "prompt-gradient", + "replicated-norm-gradient", + "hf-load-export" + ], + "gradient_max_abs": { + "A_log": 3.1377567211166024e-11, + "dt_bias": 3.069544618483633e-11, + "in_proj_qkvz.weight": 1.1714291758835316e-09, + "in_proj_ba.weight": 6.184563972055912e-11, + "conv1d.weight": 6.111804395914078e-09, + "norm.weight": 3.91155481338501e-08, + "out_proj.weight": 4.092726157978177e-10 + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-3.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-3.json new file mode 100644 index 000000000..c036f11fe --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-gdn/rank-3.json @@ -0,0 +1,32 @@ +{ + "status": "passed", + "scope": "tp4_gdn_layer0", + "rank": 3, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "cases": [ + "full-chunk-8", + "full-chunk-64", + "full-chunk-256", + "full-chunk-1024", + "prompt-gradient", + "replicated-norm-gradient", + "hf-load-export" + ], + "gradient_max_abs": { + "A_log": 8.094502845779061e-11, + "dt_bias": 7.958078640513122e-11, + "in_proj_qkvz.weight": 5.784386303275824e-10, + "in_proj_ba.weight": 3.979039320256561e-11, + "conv1d.weight": 5.093170329928398e-09, + "norm.weight": 3.91155481338501e-08, + "out_proj.weight": 3.67435859516263e-10 + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-0.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-0.json new file mode 100644 index 000000000..0fd10fafa --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-0.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 0, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 4.598405212163925e-09, + "experts.down": 7.8580342233181e-09, + "shared_expert.gate_proj.weight": 2.5890767574310303e-07, + "shared_expert.up_proj.weight": 4.3585896492004395e-07, + "shared_expert.down_proj.weight": 2.9243528842926025e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2407.167911529541, + "64": 2413.5360717773438, + "256": 1934.4639778137207, + "1024": 1961.6639614105225 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-1.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-1.json new file mode 100644 index 000000000..651abe983 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-1.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 1, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.240994945168495e-09, + "experts.down": 5.558831617236137e-09, + "shared_expert.gate_proj.weight": 3.129243850708008e-07, + "shared_expert.up_proj.weight": 3.725290298461914e-07, + "shared_expert.down_proj.weight": 4.6938657760620117e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2461.6639614105225, + "64": 2415.616035461426, + "256": 2299.391984939575, + "1024": 1961.9840383529663 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-2.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-2.json new file mode 100644 index 000000000..012e811d9 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-2.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 2, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.0227162167429924e-09, + "experts.down": 5.122274160385132e-09, + "shared_expert.gate_proj.weight": 2.337619662284851e-07, + "shared_expert.up_proj.weight": 3.9674341678619385e-07, + "shared_expert.down_proj.weight": 2.384185791015625e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2433.2799911499023, + "64": 2415.7440662384033, + "256": 2286.144018173218, + "1024": 1964.0640020370483 + } + } +} diff --git a/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-3.json b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-3.json new file mode 100644 index 000000000..44189b0c3 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-tp4-mixers-b200/tp4-moe/rank-3.json @@ -0,0 +1,45 @@ +{ + "status": "passed", + "scope": "tp4_moe_block_layer0", + "rank": 3, + "rl_kernel_commit": "9e6c4f24ce4b5e8888650e1062922711dbe565e0", + "tracked_tree_dirty": false, + "provider": "qwen3-next-shared-cuda-forward-v1", + "environment": { + "gpu": "NVIDIA B200", + "torch": "2.13.0+cu130", + "vllm": "0.30.0", + "nccl": "2.29.7" + }, + "moe": { + "cases": [ + "hf-load-export-all512experts", + "routes-replicated-8", + "full-chunk-reorder-8", + "routes-replicated-64", + "full-chunk-reorder-64", + "routes-replicated-256", + "full-chunk-reorder-256", + "routes-replicated-1024", + "full-chunk-reorder-1024", + "training-forward", + "all-parameter-gradient", + "replicated-router-gradient" + ], + "gradient_max_abs": { + "gate.weight": 1.9354047253727913e-09, + "experts.gate_up": 2.546585164964199e-09, + "experts.down": 5.296897143125534e-09, + "shared_expert.gate_proj.weight": 1.2665987014770508e-07, + "shared_expert.up_proj.weight": 2.644956111907959e-07, + "shared_expert.down_proj.weight": 2.998858690261841e-07, + "shared_expert_gate.weight": 5.513429641723633e-07 + }, + "forward_us": { + "8": 2456.5439224243164, + "64": 2408.4479808807373, + "256": 2331.712007522583, + "1024": 1955.8720588684082 + } + } +} 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/attention/deterministic_attn.py b/rl_engine/backends/cuda/attention/deterministic_attn.py index 2873c0c8e..8b6d9fc4c 100644 --- a/rl_engine/backends/cuda/attention/deterministic_attn.py +++ b/rl_engine/backends/cuda/attention/deterministic_attn.py @@ -27,8 +27,8 @@ ) from rl_engine.utils.logger import logger -_HEAD_DIM = 128 _IS_ROCM = torch.version.hip is not None +_HEAD_DIMS = (128,) if _IS_ROCM else (128, 256) _GPU_PLATFORM = "ROCm" if _IS_ROCM else "CUDA" @@ -201,8 +201,10 @@ def _validate_inputs( f"k/v shape mismatch: k={tuple(k.shape)}, v={tuple(v.shape)}, " f"expected k/v [B={b}, Hkv, Skv, D={d}]" ) - if d != _HEAD_DIM: - raise ValueError(f"head dim D must be {_HEAD_DIM}, got {d}") + if d not in _HEAD_DIMS: + raise ValueError(f"head dim D must be one of {_HEAD_DIMS}, got {d}") + if b < 1 or hq < 1 or hkv < 1: + raise ValueError("batch size and query/key head counts must be positive") if hq % hkv != 0: raise ValueError(f"Hq={hq} not divisible by Hkv={hkv} (GQA group)") if q.dtype not in (torch.float16, torch.bfloat16): @@ -211,7 +213,11 @@ def _validate_inputs( raise ValueError("q, k, v must share the same dtype") if not (q.is_cuda and k.is_cuda and v.is_cuda): raise ValueError("q, k, v must be GPU tensors") + if not (q.device == k.device == v.device): + raise ValueError("q, k, v must be on the same device") if key_padding_mask is not None: + if key_padding_mask.device != q.device: + raise ValueError("key_padding_mask must be on the input device") if key_padding_mask.dtype != torch.bool: raise ValueError("key_padding_mask must be bool") if key_padding_mask.shape != (b, skv): 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/integrations/engines/train/vllm/__init__.py b/rl_engine/integrations/engines/train/vllm/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/rl_engine/integrations/engines/train/vllm/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/rl_engine/integrations/engines/train/vllm/qwen3_next_conv.py b/rl_engine/integrations/engines/train/vllm/qwen3_next_conv.py new file mode 100644 index 000000000..2e2da8329 --- /dev/null +++ b/rl_engine/integrations/engines/train/vllm/qwen3_next_conv.py @@ -0,0 +1,133 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Differentiable, explicit-state bridge to the pinned vLLM causal-conv provider. + +Forward is the real vLLM update kernel for both prefill chunks and decode. +Backward recomputes the reference recurrence from training inputs. This module +does not install model hooks or advertise backend batch-invariance support. +""" + +from functools import lru_cache +from importlib.metadata import version + +import torch +from torch.autograd.function import once_differentiable + +from rl_engine.reference.linear_attn.causal_conv1d import CausalConv1dUpdateOp +from rl_engine.reference.linear_attn.gated_delta_rule import _validate_state_indices + + +@lru_cache(maxsize=1) +def _provider(): + if version("vllm") != "0.30.0": + raise RuntimeError("Qwen3-Next conv bridge requires the audited vLLM 0.30.0 provider") + from vllm.model_executor.layers.mamba.ops.causal_conv1d import causal_conv1d_update + + return causal_conv1d_update + + +def _lengths(x, state, weight, indices, cu_seqlens): + if x.ndim != 2 or not x.is_cuda or x.dtype != torch.bfloat16 or x.shape[1] == 0: + raise ValueError("x must be CUDA BF16 [tokens, dim] with positive dim") + if weight.shape != (x.shape[1], 4) or weight.dtype != x.dtype or weight.device != x.device: + raise ValueError("weight must be BF16 [dim, 4] on the input device") + if ( + state.ndim != 3 + or state.shape[1:] != (x.shape[1], 3) + or state.device != x.device + or state.dtype != x.dtype + ): + raise ValueError("conv state must be BF16 [blocks, dim, 3] on the input device") + if ( + cu_seqlens.ndim != 1 + or cu_seqlens.device.type != "cpu" + or cu_seqlens.dtype not in (torch.int32, torch.int64) + or cu_seqlens.numel() < 1 + ): + raise ValueError("cu_seqlens must be a nonempty CPU integer vector") + boundaries = cu_seqlens.tolist() + if boundaries[0] != 0 or boundaries[-1] != x.shape[0] or boundaries[-1] >= 2**31: + raise ValueError("cu_seqlens must cover all input tokens and fit int32") + lengths = [end - start for start, end in zip(boundaries, boundaries[1:])] + if any(length < 0 for length in lengths): + raise ValueError("cu_seqlens must be nondecreasing") + _validate_state_indices(indices, len(lengths), state.shape[0], x.device) + if bool((indices <= 0).any().item()): + raise ValueError("sequence cache indices must be positive; zero is reserved") + return lengths + + +class _CausalConvSequence(torch.autograd.Function): + @staticmethod + def forward(ctx, x, state, weight, indices, cu_seqlens): + lengths = _lengths(x, state, weight, indices, cu_seqlens) + ctx.save_for_backward(x, state, weight, indices, cu_seqlens) + new_state = state.contiguous().clone() + output = torch.zeros_like(x) + if x.shape[0]: + # The upstream varlen path subtracts (max_len - sequence_len) from + # its effective cache length, which is for speculative cache layouts. + # A width-1 ordinary cache must instead use the dense sequence path. + # Equal-length groups need neither padding nor a speculative cache. + boundaries = cu_seqlens.tolist() + for length in sorted(set(lengths) - {0}): + active = [i for i, size in enumerate(lengths) if size == length] + packed = torch.stack( + [x[boundaries[seq] : boundaries[seq + 1]].t() for seq in active] + ) + values = _provider()( + packed, + new_state, + weight.contiguous(), + activation="silu", + conv_state_indices=indices[active].to(torch.int32).contiguous(), + out=torch.empty_like(packed), + ) + for row, seq in enumerate(active): + output[boundaries[seq] : boundaries[seq + 1]] = values[row].t() + return output, new_state + + @staticmethod + @once_differentiable + def backward(ctx, grad_output, grad_state): + x, state, weight, indices, cu_seqlens = ctx.saved_tensors + with torch.enable_grad(): + inputs = [tensor.detach().requires_grad_(True) for tensor in (x, state, weight)] + x_ref, state_ref, weight_ref = inputs + pieces = [] + boundaries = cu_seqlens.tolist() + for seq, (start, end) in enumerate(zip(boundaries, boundaries[1:])): + for token in range(start, end): + out, state_ref = CausalConv1dUpdateOp()( + x_ref[token : token + 1], state_ref, weight_ref, indices[seq : seq + 1] + ) + pieces.append(out) + output = torch.cat(pieces) if pieces else x_ref[:0] + grads = torch.autograd.grad( + (output, state_ref), + inputs, + ( + torch.zeros_like(output) if grad_output is None else grad_output, + torch.zeros_like(state_ref) if grad_state is None else grad_state, + ), + allow_unused=True, + ) + return ( + *( + torch.zeros_like(value) if grad is None else grad + for value, grad in zip(inputs, grads) + ), + None, + None, + ) + + +def causal_conv_sequence(x, state, weight, indices, cu_seqlens): + """Return packed output and an independent cache; retain cache gradients. + + Qwen3-Next's width-four, bias-free SiLU convolution is the only contract. + All tensors except CPU ``cu_seqlens`` live on the input CUDA device. The + original cache and inputs are never mutated, including in inference mode. + """ + return _CausalConvSequence.apply(x, state, weight, indices, cu_seqlens) diff --git a/rl_engine/integrations/engines/train/vllm/qwen3_next_gdn.py b/rl_engine/integrations/engines/train/vllm/qwen3_next_gdn.py new file mode 100644 index 000000000..c9f136f57 --- /dev/null +++ b/rl_engine/integrations/engines/train/vllm/qwen3_next_gdn.py @@ -0,0 +1,193 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Experimental training bridge for the pinned vLLM GDN decode provider. + +Forward uses the real packed decode kernel, with a private cache copy. Backward +recomputes the mathematical recurrence from the training inputs and preserves +initial-state gradients. This is a single operator bridge, not a model installer, +prefill/decode equivalence claim, or permission to bypass vLLM's BI guard. +""" + +from functools import lru_cache +from importlib.metadata import version + +import torch +from torch.autograd.function import once_differentiable + +from rl_engine.reference.linear_attn.gated_delta_rule import ( + GatedDeltaRuleRecurrentStepOp, + _validate_state_indices, +) + + +@lru_cache(maxsize=1) +def _provider(): + if version("vllm") != "0.30.0": + raise RuntimeError("Qwen3-Next GDN bridge requires the audited vLLM 0.30.0 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 _validate(qkv, a, b, A_log, dt_bias, state, indices, num_k_heads): + if not qkv.is_cuda or qkv.ndim != 2 or qkv.dtype != torch.bfloat16: + raise ValueError("qkv must be CUDA BF16 [B, D]") + if state.ndim != 4 or state.dtype != torch.float32 or state.shape[-2:] != (128, 128): + raise ValueError("state must be FP32 [blocks, HV, 128, 128]") + 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") + batch, hv = qkv.shape[0], state.shape[1] + if hv <= 0 or hv % num_k_heads or qkv.shape[1] != 2 * num_k_heads * 128 + hv * 128: + raise ValueError("packed qkv width and state heads are inconsistent") + for name, tensor, shape, dtype in ( + ("a", a, (batch, hv), qkv.dtype), + ("b", b, (batch, hv), qkv.dtype), + ("A_log", A_log, (hv,), torch.float32), + ("dt_bias", dt_bias, (hv,), torch.float32), + ): + if tensor.shape != shape or tensor.dtype != dtype or tensor.device != qkv.device: + raise ValueError(f"{name} has an incompatible shape, dtype or device") + if state.device != qkv.device: + raise ValueError("state must share the qkv device") + _validate_state_indices(indices, batch, state.shape[0], qkv.device) + + +class _PackedDecode(torch.autograd.Function): + @staticmethod + def forward(ctx, qkv, a, b, A_log, dt_bias, state, indices, scale, num_k_heads): + _validate(qkv, a, b, A_log, dt_bias, state, indices, num_k_heads) + ctx.save_for_backward(qkv, a, b, A_log, dt_bias, state, indices) + ctx.scale, ctx.num_k_heads = float(scale), num_k_heads + new_state = state.contiguous().clone() + out = torch.zeros(qkv.shape[0], 1, state.shape[1], 128, device=qkv.device, dtype=qkv.dtype) + if qkv.shape[0]: + _provider()( + qkv.contiguous(), + a.contiguous(), + b.contiguous(), + A_log.contiguous(), + dt_bias.contiguous(), + float(scale), + new_state, + out, + indices.contiguous(), + use_qk_l2norm_in_kernel=True, + ) + return out, new_state + + @staticmethod + @once_differentiable + def backward(ctx, grad_out, grad_state): + *saved, indices = ctx.saved_tensors + with torch.enable_grad(): + inputs = [tensor.detach().requires_grad_(True) for tensor in saved] + out, state = GatedDeltaRuleRecurrentStepOp()( + *inputs, + indices, + scale=ctx.scale, + num_k_heads=ctx.num_k_heads, + ) + outputs_and_grads = [ + (value, torch.zeros_like(value) if grad is None else grad) + for value, grad in ((out, grad_out), (state, grad_state)) + if value.requires_grad + ] + grads = torch.autograd.grad( + tuple(value for value, _ in outputs_and_grads), + inputs, + tuple(grad for _, grad in outputs_and_grads), + allow_unused=True, + ) + grads = tuple( + torch.zeros_like(value) if grad is None else grad + for value, grad in zip(inputs, grads) + ) + + return (*grads, None, None, None) + + +def packed_decode_training_step( + qkv, a, b, A_log, dt_bias, state, indices, *, num_k_heads, scale=None +): + """One train-side token; pass returned state onward without detaching it. + + Supports first-order gradients only. FP32 recurrent state and Qwen3-Next + 128-wide heads are mandatory. This function never reads rollout tensors. + """ + return _PackedDecode.apply( + qkv, + a, + b, + A_log, + dt_bias, + state, + indices, + 128**-0.5 if scale is None else float(scale), + num_k_heads, + ) + + +def packed_recurrent_sequence( + qkv, a, b, A_log, dt_bias, state, indices, cu_seqlens, *, num_k_heads, scale=None +): + """Apply the decode provider to packed, variable-length sequences. + + ``cu_seqlens`` is CPU int32/int64 metadata delimiting contiguous sequences + in ``qkv[T, D]``. ``indices[B]`` maps each sequence to a distinct positive + cache slot; slot zero remains reserved by the decode provider. Returning + the entire FP32 state bank makes chunk continuation and sequence reordering + explicit. Callers must pass that state onward without detaching it for + training. No rollout cache is read and no batch-invariance capability is + registered. This intentionally slow recurrence is an integration reference, + not equivalence evidence for the upstream chunk-prefill kernel. + """ + if ( + not isinstance(cu_seqlens, torch.Tensor) + or cu_seqlens.device.type != "cpu" + or cu_seqlens.dtype not in (torch.int32, torch.int64) + or cu_seqlens.ndim != 1 + or cu_seqlens.numel() < 1 + ): + raise ValueError("cu_seqlens must be a nonempty CPU int32/int64 vector") + boundaries = cu_seqlens.tolist() + if qkv.ndim != 2 or boundaries[0] != 0 or boundaries[-1] != qkv.shape[0]: + raise ValueError("cu_seqlens must span all packed qkv tokens, starting at zero") + lengths = [end - start for start, end in zip(boundaries, boundaries[1:])] + if any(length < 0 for length in lengths): + raise ValueError("cu_seqlens must be nondecreasing") + if a.ndim != 2 or b.ndim != 2 or a.shape[0] != qkv.shape[0] or b.shape[0] != qkv.shape[0]: + raise ValueError("a and b must have one row per packed token") + _validate(qkv[:0], a[:0], b[:0], A_log, dt_bias, state, indices[:0], num_k_heads) + _validate_state_indices(indices, len(lengths), state.shape[0], qkv.device) + if bool((indices <= 0).any().item()): + raise ValueError("sequence cache indices must be positive; zero is reserved") + outputs, positions = [], [] + for position in range(max(lengths, default=0)): + active = [seq for seq, length in enumerate(lengths) if position < length] + rows = [boundaries[seq] + position for seq in active] + row_ids = torch.tensor(rows, device=qkv.device, dtype=torch.int64) + seq_ids = torch.tensor(active, device=qkv.device, dtype=torch.int64) + out, state = packed_decode_training_step( + qkv.index_select(0, row_ids), + a.index_select(0, row_ids), + b.index_select(0, row_ids), + A_log, + dt_bias, + state, + indices.index_select(0, seq_ids), + num_k_heads=num_k_heads, + scale=scale, + ) + outputs.append(out[:, 0]) + positions.extend(rows) + if not outputs: + return qkv.new_empty((0, state.shape[1], 128)), state + # Restore packed sequence order after the time-major recurrent traversal. + inverse = [0] * len(positions) + for source, destination in enumerate(positions): + inverse[destination] = source + order = torch.tensor(inverse, device=qkv.device, dtype=torch.int64) + return torch.cat(outputs).index_select(0, order), state diff --git a/rl_engine/integrations/engines/train/vllm/qwen3_next_provider.py b/rl_engine/integrations/engines/train/vllm/qwen3_next_provider.py new file mode 100644 index 000000000..af296f413 --- /dev/null +++ b/rl_engine/integrations/engines/train/vllm/qwen3_next_provider.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Shared Qwen3-Next GDN core with explicit provider and state identities. + +This is the convolution/recurrence/gated-norm boundary, after the projections +and before the output projection. It is not a full attention or model adapter. +The same call can be made under inference_mode or with training autograd. +""" + +from dataclasses import asdict, dataclass + +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_conv import causal_conv_sequence +from rl_engine.integrations.engines.train.vllm.qwen3_next_gdn import packed_recurrent_sequence + + +@dataclass(frozen=True) +class GDNProviderConfig: + """The only implemented profile; other providers cannot silently substitute.""" + + version: str = "qwen3-next-gdn-core-v1" + convolution: str = "vllm-0.30.0-causal-conv-update-equal-length-groups" + recurrence: str = "vllm-0.30.0-packed-recurrent-decode" + gated_norm: str = "rl-engine-cuda-rmsnorm-gated" + recurrent_dtype: str = "float32" + activation_dtype: str = "bfloat16" + head_dim: int = 128 + + def validate(self): + if self != GDNProviderConfig(): + raise ValueError("Unsupported GDN provider configuration") + + def identity(self): + self.validate() + return asdict(self) + + +@dataclass(frozen=True) +class GDNState: + convolution: torch.Tensor + recurrent: torch.Tensor + + +def shared_gdn_core( + qkv, + a, + b, + z, + A_log, + dt_bias, + conv_weight, + norm_weight, + state, + indices, + cu_seqlens, + *, + config, + num_k_heads, + eps=1e-6, +): + """Independently recompute packed training inputs or execute rollout inputs. + + qkv is laid out as contiguous Q, K, V groups. z is [tokens, value_heads, 128]. + Both callers must explicitly supply the same GDNProviderConfig. State is + returned without detachment so a response loss can differentiate through + the prompt. No global vLLM batch-invariance capability is changed. + """ + if not isinstance(config, GDNProviderConfig): + raise ValueError("An explicit GDNProviderConfig is required") + config.validate() + if not isinstance(state, GDNState): + raise ValueError("An explicit convolution and recurrent GDNState is required") + if state.recurrent.ndim != 4: + raise ValueError("Recurrent state must have four dimensions") + if ( + z.shape != (qkv.shape[0], state.recurrent.shape[1], 128) + or z.dtype != torch.bfloat16 + or z.device != qkv.device + ): + raise ValueError("z must be BF16 [tokens, value_heads, 128] on the input device") + if ( + norm_weight.shape != (128,) + or norm_weight.dtype != torch.bfloat16 + or norm_weight.device != qkv.device + ): + raise ValueError("Gated norm weight must be BF16 [128] on the input device") + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + norm = Qwen3NextRMSNormGatedCudaOp() + + convolved, conv_state = causal_conv_sequence( + qkv, state.convolution, conv_weight, indices, cu_seqlens + ) + recurrent, recurrent_state = packed_recurrent_sequence( + convolved, + a, + b, + A_log, + dt_bias, + state.recurrent, + indices, + cu_seqlens, + num_k_heads=num_k_heads, + ) + normed = norm(recurrent.reshape(-1, 128), norm_weight, z.reshape(-1, 128), eps=eps).reshape_as( + z + ) + return normed, GDNState(conv_state, recurrent_state) diff --git a/rl_engine/models/qwen3_next/__init__.py b/rl_engine/models/qwen3_next/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/rl_engine/models/qwen3_next/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/rl_engine/models/qwen3_next/qwen3_next_forward.py b/rl_engine/models/qwen3_next/qwen3_next_forward.py new file mode 100644 index 000000000..d5b0f0656 --- /dev/null +++ b/rl_engine/models/qwen3_next/qwen3_next_forward.py @@ -0,0 +1,364 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Explicit CUDA forward primitives for the Qwen3-Next shared TP4 profile. + +These functions do not install framework hooks or advertise engine capabilities. +The caller owns parameter sharding, TP collectives, positions and cache identity. +Weights remain BF16; only the router explicitly computes in FP32, as in VIME's +Qwen3-Next configuration. No CPU or alternate-kernel fallback is permitted. +""" + +from dataclasses import dataclass +from functools import lru_cache +from importlib.metadata import version + +import torch +import torch.nn.functional as F +from torch.autograd.function import once_differentiable + +FORWARD_PROVIDER_ID = "qwen3-next-shared-cuda-forward-v1" + + +@lru_cache(maxsize=1) +def _vllm_linear(): + if version("vllm") != "0.30.0": + raise RuntimeError("Shared Qwen3-Next linear requires pinned vLLM 0.30.0") + from vllm.model_executor.determinism.batch_invariant import linear_batch_invariant + + return linear_batch_invariant + + +# The batch-invariant tile vLLM itself selects under VLLM_BATCH_INVARIANT=1 +# (fused_moe.get_default_config), passed explicitly so no environment variable, +# tuned-config file or token count can change it. +_GROUPED_CONFIG = { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 8, + "SPLIT_K": 1, +} + + +@lru_cache(maxsize=1) +def _vllm_grouped(): + if version("vllm") != "0.30.0": + raise RuntimeError("Shared Qwen3-Next experts require pinned vLLM 0.30.0") + from vllm.model_executor.layers.fused_moe.fused_moe import ( + _prepare_expert_assignment, + dispatch_fused_moe_kernel, + ) + from vllm.model_executor.layers.fused_moe.utils import resolve_moe_use_td + + return _prepare_expert_assignment, dispatch_fused_moe_kernel, resolve_moe_use_td + + +def grouped_route_linear(rows, weight, route_indices): + """Per-route ``rows @ weight[expert].T`` through vLLM's fused MoE Triton kernel. + + ``rows`` is [T, K] (one row per token, shared by its ten routes) or + [T*10, K] (one row per route); the result is [T, 10, N] in route order. + No routed weight is applied and nothing is summed: every output row has + exactly one writer, and its K loop has the fixed order of the fixed tile. + """ + prepare, dispatch, use_td = _vllm_grouped() + if use_td(): + raise RuntimeError("VLLM_TRITON_USE_TD changes the MoE kernel; strict mode requires 0") + import triton.language as tl + + tokens, top_k = route_indices.shape + out = rows.new_empty((tokens, top_k, weight.shape[1])) + if tokens == 0: + return out + indices = route_indices.to(torch.int32).contiguous() + sorted_ids, expert_ids, padded = prepare( + indices, _GROUPED_CONFIG, tokens, top_k, weight.shape[0], None, ignore_invalid_experts=True + ) + dispatch( + rows.contiguous(), + weight, + out, + None, + None, + None, + None, + sorted_ids, + expert_ids, + padded, + False, + top_k if rows.shape[0] == tokens else 1, + _GROUPED_CONFIG, + compute_type=tl.bfloat16, + use_fp8_w8a8=False, + use_int8_w8a8=False, + use_int8_w8a16=False, + use_int4_w4a16=False, + per_channel_quant=False, + ) + return out + + +def _cuda_tensors(*tensors): + if torch.version.hip is not None or any(not value.is_cuda for value in tensors): + raise ValueError("Shared Qwen3-Next forward requires CUDA tensors; no CPU fallback") + if any(value.device != tensors[0].device for value in tensors[1:]): + raise ValueError("Shared Qwen3-Next tensors must be on the same device") + + +class _SharedLinear(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight, bias): + ctx.save_for_backward(x, weight) + ctx.has_bias = bias is not None + if x.shape[0] == 0: + return x.new_empty((0, weight.shape[0])) + with torch.cuda.device(x.device): + return _vllm_linear()(x, weight, bias) + + @staticmethod + @once_differentiable + def backward(ctx, grad_output): + x, weight = ctx.saved_tensors + if x.shape[0] == 0: + return ( + torch.empty_like(x) if ctx.needs_input_grad[0] else None, + torch.zeros_like(weight) if ctx.needs_input_grad[1] else None, + weight.new_zeros(weight.shape[0]) if ctx.has_bias else None, + ) + grad = grad_output.contiguous() + with torch.cuda.device(x.device): + linear = _vllm_linear() + dx = linear(grad, weight.t()) if ctx.needs_input_grad[0] else None + dw = linear(grad.t(), x.t()) if ctx.needs_input_grad[1] else None + db = grad.sum(dim=0) if ctx.has_bias and ctx.needs_input_grad[2] else None + return dx, dw, db + + +def shared_linear(x, weight, bias=None): + """Pinned vLLM deterministic GEMM with an explicit first-order VJP. + + x is [..., K], weight is [N, K], output is [..., N]. Operands have one + common BF16 or FP32 dtype. Empty token dimensions preserve autograd. + Backward executes dX=dY@W and dW=dY.T@X through the same fixed provider. + """ + _cuda_tensors(x, weight, *(() if bias is None else (bias,))) + if x.ndim < 2 or weight.ndim != 2 or x.shape[-1] != weight.shape[1]: + raise ValueError("Shared linear requires x[..., K] and weight[N, K]") + if x.dtype not in (torch.bfloat16, torch.float32) or weight.dtype != x.dtype: + raise ValueError("Shared linear operands must have one common BF16 or FP32 dtype") + if x.dtype == torch.float32: + import triton + + if triton.knobs.language.fp32_default != "ieee": + raise RuntimeError( + "FP32 shared linear requires TRITON_F32_DEFAULT=ieee at process startup" + ) + if min(weight.shape) < 1: + raise ValueError("Shared linear weight dimensions must be positive") + if bias is not None and (bias.shape != (weight.shape[0],) or bias.dtype != x.dtype): + raise ValueError("Shared linear bias must be [N] with the input dtype") + out = _SharedLinear.apply(x.reshape(-1, x.shape[-1]), weight, bias) + return out.reshape(*x.shape[:-1], weight.shape[0]) + + +def shared_router(x, weight): + """Compute the 512-expert FP32 router from BF16 inputs/parameter storage.""" + _cuda_tensors(x, weight) + if x.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16: + raise ValueError("Qwen3-Next router input and stored weight must be BF16") + if x.ndim != 2 or weight.shape != (512, x.shape[-1]): + raise ValueError("Qwen3-Next router requires x[T, H] and weight[512, H]") + return shared_linear(x.float(), weight.float()) + + +@dataclass(frozen=True) +class Routing: + """Descending probability routes; equal scores prefer smaller expert IDs.""" + + weights: torch.Tensor + indices: torch.Tensor + + +def fixed_order_row_sum(values): + """Sum the last dimension by a pairwise tree of elementwise adds. + + ``torch.sum`` picks its reduction order from the tensor shape, so the same + row summed inside an 8-row and a 32-row tensor can differ by one ulp unless + vLLM's batch-invariant aten overrides are installed. Qualification runs + execute with ``VLLM_BATCH_INVARIANT=0``, where those overrides are absent. + Elementwise adds have no such freedom: every row is summed in one fixed + tree whatever the row count. Odd widths are padded with exact zeros. + """ + while values.shape[-1] > 1: + if values.shape[-1] % 2: + values = torch.cat((values, values.new_zeros((*values.shape[:-1], 1))), dim=-1) + values = values[..., 0::2] + values[..., 1::2] + return values + + +def fixed_order_softmax(logits): + """Row softmax from a max subtraction, exp and one fixed-order row sum.""" + shifted = logits - torch.amax(logits, dim=-1, keepdim=True) + exponentials = torch.exp(shifted) + return exponentials / fixed_order_row_sum(exponentials) + + +def stable_top10_routes(router_logits): + """Select ten distinct experts and normalize their FP32 probabilities. + + The full 512-way softmax precedes selection, preserving Qwen's normalized + top-k arithmetic. Stable sort fixes the otherwise ambiguous tie boundary. + Indices are discrete; gradients flow through the selected probabilities. + Both row sums use the fixed-order tree, so a token's route weights do not + depend on how many other tokens share its forward pass. + """ + _cuda_tensors(router_logits) + if router_logits.ndim != 2 or router_logits.shape[1] != 512: + raise ValueError("Qwen3-Next routing requires logits[T, 512]") + if router_logits.dtype != torch.float32: + raise ValueError("Qwen3-Next router logits must be FP32") + if not torch.isfinite(router_logits).all(): + raise ValueError("Qwen3-Next router logits must be finite") + probabilities = fixed_order_softmax(router_logits) + indices = torch.argsort(router_logits, dim=-1, descending=True, stable=True)[:, :10] + selected = probabilities.gather(1, indices) + return Routing(selected / fixed_order_row_sum(selected), indices) + + +def combine_routes(expert_outputs, route_weights): + """Combine [T, 10, H] BF16 outputs by ten ordered FP32 multiply-adds. + + There are no atomic sums. Each route's product and each addition is a + separate eager operation; the single BF16 cast follows the final route. + """ + _cuda_tensors(expert_outputs, route_weights) + if expert_outputs.ndim != 3 or expert_outputs.shape[1] != 10: + raise ValueError("Qwen3-Next expert outputs must have shape [T, 10, H]") + if route_weights.shape != expert_outputs.shape[:2]: + raise ValueError("Route weights must have shape [T, 10]") + if expert_outputs.dtype != torch.bfloat16 or route_weights.dtype != torch.float32: + raise ValueError("Combine requires BF16 expert outputs and FP32 route weights") + output = expert_outputs[:, 0].float() * route_weights[:, 0, None] + for slot in range(1, 10): + output = output + expert_outputs[:, slot].float() * route_weights[:, slot, None] + return output.to(torch.bfloat16) + + +def _expert_activation(gate_up_rows): + """SwiGLU in FP32 from the BF16 [gate, up] projection; returns its parts too.""" + gate, up = gate_up_rows.float().chunk(2, dim=-1) + silu_gate = F.silu(gate) + return silu_gate * up, silu_gate, gate, up + + +class _RoutedExperts(torch.autograd.Function): + """Grouped expert forward; per-expert backward into one dense gradient buffer.""" + + @staticmethod + def forward(ctx, x, gate_up, down, route_indices): + ctx.save_for_backward(x, gate_up, down, route_indices) + with torch.cuda.device(x.device): + projected = grouped_route_linear(x, gate_up, route_indices) + activated = _expert_activation(projected)[0] + activated = activated.to(x.dtype).reshape(-1, activated.shape[-1]) + return grouped_route_linear(activated, down, route_indices) + + @staticmethod + @once_differentiable + def backward(ctx, grad_output): + x, gate_up, down, route_indices = ctx.saved_tensors + dx = torch.zeros_like(x) if ctx.needs_input_grad[0] else None + dgate_up = torch.zeros_like(gate_up) if ctx.needs_input_grad[1] else None + ddown = torch.zeros_like(down) if ctx.needs_input_grad[2] else None + top_k = route_indices.shape[1] + slots_by_expert = [[] for _ in range(gate_up.shape[0])] + for slot, expert in enumerate(route_indices.flatten().tolist()): + slots_by_expert[expert].append(slot) + grad_flat = grad_output.reshape(-1, x.shape[1]) + with torch.cuda.device(x.device): + linear = _vllm_linear() + projected = grouped_route_linear(x, gate_up, route_indices) + projected = projected.reshape(-1, projected.shape[-1]) + for expert, slots in enumerate(slots_by_expert): + if not slots: + continue + slots = torch.tensor(slots, dtype=torch.long, device=x.device) + tokens = slots // top_k + activated, silu_gate, gate32, up32 = _expert_activation( + projected.index_select(0, slots) + ) + activated = activated.to(x.dtype) + grad = grad_flat.index_select(0, slots) + if ddown is not None: + ddown[expert].copy_(linear(grad.t(), activated.t())) + if dx is None and dgate_up is None: + continue + rows = x.index_select(0, tokens) + dactivated = linear(grad, down[expert].t()).float() + sigmoid = gate32.sigmoid() + dgate = dactivated * up32 * (sigmoid + gate32 * sigmoid * (1 - sigmoid)) + dup = dactivated * silu_gate + dprojected = torch.cat((dgate, dup), dim=-1).to(x.dtype) + if dgate_up is not None: + dgate_up[expert].copy_(linear(dprojected.t(), rows.t())) + if dx is not None: + # Top-k has distinct experts, so tokens are unique within this + # launch. Ascending expert order fixes each token's sum order. + partial = linear(dprojected, gate_up[expert].t()) + dx.index_copy_(0, tokens, dx.index_select(0, tokens) + partial) + return dx, dgate_up, ddown, None + + +def shared_moe(x, router_weight, gate_up_weights, down_weights): + """Evaluate TP-local routed experts; caller reduces the returned output. + + Weights are [512, 2*I_local, H] (gate then up) and [512, H, I_local]. + All experts belong to EP=1. Each (token, route) output row has exactly one + writer and the ten routes are combined in route order. Shared-expert and + TP-reduction ownership remains with the model adapter. The forward has no + host synchronisation; the backward builds per-expert row lists on the host. + """ + _cuda_tensors(x, router_weight, gate_up_weights, down_weights) + if x.ndim != 2 or x.dtype != torch.bfloat16: + raise ValueError("Qwen3-Next MoE input must be BF16 [T, H]") + if gate_up_weights.ndim != 3 or down_weights.ndim != 3: + raise ValueError("Qwen3-Next expert weights must have three dimensions") + local_intermediate = down_weights.shape[-1] + if local_intermediate < 1: + raise ValueError("Expert intermediate dimension must be positive") + if gate_up_weights.shape != (512, 2 * local_intermediate, x.shape[1]): + raise ValueError("gate_up_weights must have shape [512, 2*I_local, H]") + if down_weights.shape != (512, x.shape[1], local_intermediate): + raise ValueError("down_weights must have shape [512, H, I_local]") + if gate_up_weights.dtype != x.dtype or down_weights.dtype != x.dtype: + raise ValueError("Qwen3-Next expert weights must remain BF16") + routes = stable_top10_routes(shared_router(x, router_weight)) + outputs = _RoutedExperts.apply(x, gate_up_weights, down_weights, routes.indices) + return combine_routes(outputs, routes.weights), routes + + +@lru_cache(maxsize=1) +def _attention_op(): + from rl_engine.backends.cuda.attention.deterministic_attn import DeterministicAttentionOp + + return DeterministicAttentionOp() + + +def shared_attention(q, k, v): + """Qwen TP4 GQA core: BF16 [B,4,Sq,256] over [B,1,Skv,256]. + + Queries must be the final Sq positions of the supplied KV prefix. This + covers full prefill, any contiguous prefill chunk, and one-token decode. + KV tensors can carry the prompt's gradient graph; no state is detached. + """ + _cuda_tensors(q, k, v) + if q.ndim != 4 or k.ndim != 4 or v.ndim != 4: + raise ValueError("Qwen3-Next attention requires four-dimensional Q/K/V") + if q.shape[1] != 4 or k.shape[1] != 1 or q.shape[-1] != 256: + raise ValueError("Qwen3-Next TP4 attention requires Q heads=4, KV heads=1, D=256") + if q.dtype != torch.bfloat16 or k.dtype != q.dtype or v.dtype != q.dtype: + raise ValueError("Qwen3-Next attention inputs must be BF16") + if q.shape[2] > k.shape[2]: + raise ValueError("Query tokens must fit within the supplied KV prefix") + return _attention_op()(q, k, v, causal=True, scale=1 / 16) diff --git a/rl_engine/models/qwen3_next/qwen3_next_tp.py b/rl_engine/models/qwen3_next/qwen3_next_tp.py new file mode 100644 index 000000000..876d373ff --- /dev/null +++ b/rl_engine/models/qwen3_next/qwen3_next_tp.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Differentiable TP4 boundaries and parameter metadata shared by Qwen3-Next blocks. + +The forward copy and reduction follow Megatron's column/row-parallel autograd +pair. Reductions go through ``collective_for_group`` so a strict deployment +can pin the collective algorithm; nothing here selects one implicitly. +""" + +import torch +import torch.distributed as dist + + +def _reduce(value, group): + from rl_engine.distributed.algorithms.collectives import collective_for_group + + collective = collective_for_group(group, min_size_bytes=value.numel() * value.element_size()) + return collective.all_reduce(value.contiguous()) + + +class _CopyToTP(torch.autograd.Function): + @staticmethod + def forward(ctx, value, group): + ctx.group = group + return value + + @staticmethod + def backward(ctx, gradient): + return _reduce(gradient, ctx.group), None + + +class _ReduceFromTP(torch.autograd.Function): + @staticmethod + def forward(ctx, value, group): + return _reduce(value, group) + + @staticmethod + def backward(ctx, gradient): + return gradient, None + + +def _rank(rank): + if isinstance(rank, bool) or not isinstance(rank, int) or rank not in range(4): + raise ValueError("TP4 rank must be an integer in [0, 4)") + + +def _weight(value, shape, name): + if value.dtype != torch.bfloat16 or value.shape != shape: + raise ValueError(f"{name} must be BF16 {shape}") + + +def _tp_group(group): + if not dist.is_initialized() or dist.get_world_size(group) != 4: + raise ValueError("Qwen3-Next blocks require an initialized four-rank TP group") + return dist.group.WORLD if group is None else group + + +def _parallel_parameter(parameter, dimension, *, stride=1, duplicate=False): + # MCore uses these attributes when computing the global optimizer gradient + # norm. A replicated copy updates normally but must not count twice. + parameter.tensor_model_parallel = True + parameter.partition_dim = dimension + parameter.partition_stride = stride + parameter.shared = duplicate diff --git a/rl_engine/models/qwen3_next/qwen3_next_tp_blocks.py b/rl_engine/models/qwen3_next/qwen3_next_tp_blocks.py new file mode 100644 index 000000000..7062fe527 --- /dev/null +++ b/rl_engine/models/qwen3_next/qwen3_next_tp_blocks.py @@ -0,0 +1,375 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Official Qwen3-Next TP4 full-attention and MoE blocks for shared adapters. + +These blocks own the local arithmetic and differentiable TP boundaries. They +do not register an engine backend or advertise full-model acceptance. +""" + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from functools import lru_cache + +import torch +import torch.distributed as dist +import torch.nn.functional as F +from torch import nn + +from rl_engine.models.qwen3_next.qwen3_next_forward import ( + shared_attention, + shared_linear, + shared_moe, +) +from rl_engine.models.qwen3_next.qwen3_next_tp import ( + _CopyToTP, + _parallel_parameter, + _rank, + _ReduceFromTP, + _tp_group, + _weight, +) +from rl_engine.validation.common.tensor_identity import tensor_bitwise_equal + +ATTENTION_HF_SHAPES = { + "q_proj.weight": (8192, 2048), + "k_proj.weight": (512, 2048), + "v_proj.weight": (512, 2048), + "o_proj.weight": (2048, 4096), + "q_norm.weight": (256,), + "k_norm.weight": (256,), +} + + +def shard_attention_weights(weights: Mapping[str, torch.Tensor], rank: int): + """Preserve Q/gate interleaving and replicate each KV head on a rank pair.""" + _rank(rank) + if set(weights) != set(ATTENTION_HF_SHAPES): + raise ValueError("Attention weight names do not match the official HF contract") + result = {} + for name, shape in ATTENTION_HF_SHAPES.items(): + value = weights[name] + _weight(value, shape, name) + if name.startswith(("q_norm.", "k_norm.")): + part = value + elif name in ("k_proj.weight", "v_proj.weight"): + part = value.chunk(2, dim=0)[rank // 2] + else: + part = value.chunk(4, dim=1 if name == "o_proj.weight" else 0)[rank] + result[name] = part.contiguous() + return result + + +def assemble_attention_weights(shards: Sequence[Mapping[str, torch.Tensor]]): + """Invert four local HF mappings and reject divergent replicated storage.""" + if len(shards) != 4 or any(set(shard) != set(ATTENTION_HF_SHAPES) for shard in shards): + raise ValueError("Exactly four complete attention HF shards are required") + result = {} + for name, shape in ATTENTION_HF_SHAPES.items(): + values = [shard[name] for shard in shards] + if name.startswith(("q_norm.", "k_norm.")): + for value in values: + _weight(value, shape, name) + if any(not tensor_bitwise_equal(value, values[0]) for value in values[1:]): + raise ValueError(f"Replicated attention norm differs: {name}") + result[name] = values[0] + elif name in ("k_proj.weight", "v_proj.weight"): + for value in values: + _weight(value, (256, 2048), name) + if not tensor_bitwise_equal(values[0], values[1]) or not tensor_bitwise_equal( + values[2], values[3] + ): + raise ValueError(f"Replicated KV head differs within its rank pair: {name}") + result[name] = torch.cat((values[0], values[2]), dim=0) + else: + axis = 1 if name == "o_proj.weight" else 0 + local_shape = list(shape) + local_shape[axis] //= 4 + for value in values: + _weight(value, tuple(local_shape), name) + result[name] = torch.cat(values, dim=axis) + return result + + +@lru_cache(maxsize=4) +def _kv_replica_groups(group): + ranks = dist.get_process_group_ranks(group) + # Every TP rank creates both groups in the same order. Only members use a + # group's collective, and the rank pairs follow the official KV layout. + return tuple( + dist.new_group(ranks[start : start + 2], use_local_synchronization=True) for start in (0, 2) + ) + + +def _norm_module(device): + module = nn.Module() + module.register_parameter( + "weight", nn.Parameter(torch.zeros(256, device=device, dtype=torch.bfloat16)) + ) + return module + + +@dataclass(frozen=True) +class KVState: + """Per-slot [1,1,length,256] KV tensors; length is the next absolute position.""" + + keys: tuple[torch.Tensor, ...] + values: tuple[torch.Tensor, ...] + + +def _packed_sequences(hidden, state, indices, cu_seqlens): + if ( + hidden.ndim != 2 + or hidden.shape[1] != 2048 + or hidden.dtype != torch.bfloat16 + or not hidden.is_cuda + ): + raise ValueError("Attention hidden must be CUDA BF16 [tokens, 2048]") + if hidden.shape[0] < 1: + raise ValueError("Attention requires at least one active token") + if ( + not isinstance(state, KVState) + or len(state.keys) < 2 + or len(state.keys) != len(state.values) + ): + raise ValueError("KVState requires a reserved slot and matching key/value slots") + if ( + indices.ndim != 1 + or indices.dtype not in (torch.int32, torch.int64) + or indices.device != hidden.device + ): + raise ValueError("State indices must be an integer vector on the input device") + if ( + cu_seqlens.device.type != "cpu" + or cu_seqlens.ndim != 1 + or cu_seqlens.dtype not in (torch.int32, torch.int64) + ): + raise ValueError("cu_seqlens must be a CPU integer vector") + slots, ends = indices.tolist(), cu_seqlens.tolist() + if len(ends) != len(slots) + 1 or ends[0] != 0 or ends[-1] != hidden.shape[0]: + raise ValueError("cu_seqlens must describe all packed tokens") + if any(left > right for left, right in zip(ends, ends[1:])): + raise ValueError("cu_seqlens must be nondecreasing") + if len(set(slots)) != len(slots) or any(slot <= 0 or slot >= len(state.keys) for slot in slots): + raise ValueError("Active state slots must be unique and exclude reserved slot zero") + for key, value in zip(state.keys, state.values): + if ( + key.ndim != 4 + or key.shape[:2] != (1, 1) + or key.shape[-1] != 256 + or key.shape != value.shape + ): + raise ValueError("KV slots must have matching [1,1,length,256] tensors") + if ( + key.dtype != hidden.dtype + or value.dtype != hidden.dtype + or key.device != hidden.device + or value.device != hidden.device + ): + raise ValueError("KV slots must have the input device and BF16 dtype") + return slots, ends + + +class TP4FullAttention(nn.Module): + """Four local query heads and one pair-replicated KV head, with explicit state.""" + + def __init__(self, *, group, device, eps=1e-6, rope_theta=10000000): + super().__init__() + self.group = _tp_group(group) + self.rank = dist.get_rank(self.group) + self.kv_group = _kv_replica_groups(self.group)[self.rank // 2] + if eps != 1e-6 or rope_theta != 10000000: + raise ValueError("Full attention requires official epsilon and RoPE theta") + self.eps = eps + factory = {"device": device, "dtype": torch.bfloat16} + self.q_proj = nn.Linear(2048, 2048, bias=False, **factory) + self.k_proj = nn.Linear(2048, 256, bias=False, **factory) + self.v_proj = nn.Linear(2048, 256, bias=False, **factory) + self.o_proj = nn.Linear(1024, 2048, bias=False, **factory) + self.q_norm, self.k_norm = _norm_module(device), _norm_module(device) + _parallel_parameter(self.q_proj.weight, 0) + _parallel_parameter(self.o_proj.weight, 1) + _parallel_parameter(self.k_proj.weight, 0, duplicate=self.rank % 2 == 1) + _parallel_parameter(self.v_proj.weight, 0, duplicate=self.rank % 2 == 1) + # Kept as a Python float, not a buffer: Megatron's Float16Module casts every + # floating-point buffer of the actor to BF16, and BF16-rounded inverse + # frequencies rotate keys differently from the rollout engine's FP32 ones. + self.rope_theta = float(rope_theta) + + def load_hf_weights(self, weights): + self.load_state_dict(shard_attention_weights(weights, self.rank), strict=True) + + def export_local_hf_weights(self): + return self.state_dict() + + def initial_state(self, blocks): + if isinstance(blocks, bool) or not isinstance(blocks, int) or blocks < 2: + raise ValueError("KV state requires a reserved slot and an active slot") + return KVState( + tuple(self.q_proj.weight.new_empty((1, 1, 0, 256)) for _ in range(blocks)), + tuple(self.q_proj.weight.new_empty((1, 1, 0, 256)) for _ in range(blocks)), + ) + + @staticmethod + def inverse_frequencies(rope_theta, device): + """FP32 inverse frequencies of the 64-wide rotary quarter, computed on demand.""" + exponents = torch.arange(0, 64, 2, dtype=torch.float32, device=device) / 64 + return 1.0 / (float(rope_theta) ** exponents) + + def _rotate(self, value, positions): + inv_freq = TP4FullAttention.inverse_frequencies(self.rope_theta, value.device) + angles = positions.float()[:, None] * inv_freq[None, :] + angles = torch.cat((angles, angles), dim=-1) + cos, sin = angles.cos().to(value.dtype)[:, None], angles.sin().to(value.dtype)[:, None] + rotary, tail = value[..., :64], value[..., 64:] + rotated = torch.cat((-rotary[..., 32:], rotary[..., :32]), dim=-1) + return torch.cat((rotary * cos + rotated * sin, tail), dim=-1) + + def forward(self, hidden, state, indices, cu_seqlens): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp + + slots, ends = _packed_sequences(hidden, state, indices, cu_seqlens) + copied = _CopyToTP.apply(hidden, self.group) + query, gate = shared_linear(copied, self.q_proj.weight).reshape(-1, 4, 512).chunk(2, dim=-1) + key = shared_linear(copied, _CopyToTP.apply(self.k_proj.weight, self.kv_group)).reshape( + -1, 1, 256 + ) + value = shared_linear(copied, _CopyToTP.apply(self.v_proj.weight, self.kv_group)).reshape( + -1, 1, 256 + ) + norm = Qwen3NextRMSNormCudaOp() + query = norm(query, _CopyToTP.apply(self.q_norm.weight, self.group), eps=self.eps) + key = norm(key, _CopyToTP.apply(self.k_norm.weight, self.group), eps=self.eps) + next_keys, next_values = list(state.keys), list(state.values) + outputs = [] + for slot, begin, end in zip(slots, ends[:-1], ends[1:]): + if begin == end: + continue + offset = state.keys[slot].shape[2] + if offset + end - begin > 262144: + raise ValueError("KV continuation exceeds the official maximum position") + positions = torch.arange(offset, offset + end - begin, device=hidden.device) + q = self._rotate(query[begin:end], positions).transpose(0, 1).unsqueeze(0) + k = self._rotate(key[begin:end], positions).transpose(0, 1).unsqueeze(0) + v = value[begin:end].transpose(0, 1).unsqueeze(0) + next_keys[slot] = torch.cat((state.keys[slot], k), dim=2) + next_values[slot] = torch.cat((state.values[slot], v), dim=2) + out = shared_attention(q, next_keys[slot], next_values[slot]) + outputs.append(out.squeeze(0).transpose(0, 1)) + attended = torch.cat(outputs, dim=0) * gate.sigmoid() + local = shared_linear(attended.reshape(-1, 1024), self.o_proj.weight) + return _ReduceFromTP.apply(local, self.group), KVState(tuple(next_keys), tuple(next_values)) + + +def moe_hf_shapes(): + result = { + "gate.weight": (512, 2048), + "shared_expert_gate.weight": (1, 2048), + } + for prefix in ("shared_expert", *(f"experts.{expert}" for expert in range(512))): + result[f"{prefix}.gate_proj.weight"] = (512, 2048) + result[f"{prefix}.up_proj.weight"] = (512, 2048) + result[f"{prefix}.down_proj.weight"] = (2048, 512) + return result + + +def shard_moe_weight(name, value, rank): + """Shard one official HF tensor without materializing all expert weights.""" + _rank(rank) + shapes = moe_hf_shapes() + if name not in shapes: + raise ValueError(f"Unknown Qwen3-Next MoE weight: {name}") + _weight(value, shapes[name], name) + if name in ("gate.weight", "shared_expert_gate.weight"): + return value + return value.chunk(4, dim=1 if name.endswith("down_proj.weight") else 0)[rank].contiguous() + + +def assemble_moe_weight(name, values): + if len(values) != 4: + raise ValueError("Exactly four MoE weight shards are required") + shapes = moe_hf_shapes() + if name not in shapes: + raise ValueError(f"Unknown Qwen3-Next MoE weight: {name}") + if name in ("gate.weight", "shared_expert_gate.weight"): + for value in values: + _weight(value, shapes[name], name) + if any(not tensor_bitwise_equal(value, values[0]) for value in values[1:]): + raise ValueError(f"Replicated MoE weight differs across TP: {name}") + return values[0] + axis = 1 if name.endswith("down_proj.weight") else 0 + shape = list(shapes[name]) + shape[axis] //= 4 + for value in values: + _weight(value, tuple(shape), name) + return torch.cat(tuple(values), dim=axis) + + +class TP4MoE(nn.Module): + """All 512 experts at EP1 with intermediate width 128 per TP rank.""" + + def __init__(self, *, group, device): + super().__init__() + self.group = _tp_group(group) + self.rank = dist.get_rank(self.group) + factory = {"device": device, "dtype": torch.bfloat16} + self.gate = nn.Linear(2048, 512, bias=False, **factory) + self.experts = nn.Module() + self.experts.register_parameter( + "gate_up", nn.Parameter(torch.empty(512, 256, 2048, **factory)) + ) + self.experts.register_parameter( + "down", nn.Parameter(torch.empty(512, 2048, 128, **factory)) + ) + self.shared_expert = nn.Module() + self.shared_expert.gate_proj = nn.Linear(2048, 128, bias=False, **factory) + self.shared_expert.up_proj = nn.Linear(2048, 128, bias=False, **factory) + self.shared_expert.down_proj = nn.Linear(128, 2048, bias=False, **factory) + self.shared_expert_gate = nn.Linear(2048, 1, bias=False, **factory) + _parallel_parameter(self.experts.gate_up, 1, stride=2) + _parallel_parameter(self.experts.down, 2) + _parallel_parameter(self.shared_expert.gate_proj.weight, 0) + _parallel_parameter(self.shared_expert.up_proj.weight, 0) + _parallel_parameter(self.shared_expert.down_proj.weight, 1) + + def _hf_tensor(self, name): + if name.startswith("experts."): + _, expert, projection, _ = name.split(".") + expert = int(expert) + if projection == "down_proj": + return self.experts.down[expert] + offset = 0 if projection == "gate_proj" else 128 + return self.experts.gate_up[expert, offset : offset + 128] + obj = self + for part in name.split("."): + obj = getattr(obj, part) + return obj + + def load_hf_weights(self, weights): + if set(weights) != set(moe_hf_shapes()): + raise ValueError("MoE weight names do not match the official HF contract") + with torch.no_grad(): + for name in moe_hf_shapes(): + self._hf_tensor(name).copy_(shard_moe_weight(name, weights[name], self.rank)) + + def export_local_hf_weights(self): + return {name: self._hf_tensor(name).detach() for name in moe_hf_shapes()} + + def forward(self, hidden): + if hidden.ndim != 2 or hidden.shape[1] != 2048 or hidden.dtype != torch.bfloat16: + raise ValueError("MoE hidden must be BF16 [tokens, 2048]") + copied = _CopyToTP.apply(hidden, self.group) + routed, _ = shared_moe( + copied, + _CopyToTP.apply(self.gate.weight, self.group), + self.experts.gate_up, + self.experts.down, + ) + gate = shared_linear(copied, self.shared_expert.gate_proj.weight) + up = shared_linear(copied, self.shared_expert.up_proj.weight) + activated = (F.silu(gate.float()) * up.float()).to(hidden.dtype) + shared = shared_linear(activated, self.shared_expert.down_proj.weight) + shared_gate = shared_linear( + copied, _CopyToTP.apply(self.shared_expert_gate.weight, self.group) + ).sigmoid() + return _ReduceFromTP.apply(routed + shared * shared_gate, self.group) diff --git a/rl_engine/models/qwen3_next/qwen3_next_tp_gdn.py b/rl_engine/models/qwen3_next/qwen3_next_tp_gdn.py new file mode 100644 index 000000000..07abaea20 --- /dev/null +++ b/rl_engine/models/qwen3_next/qwen3_next_tp_gdn.py @@ -0,0 +1,163 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""TP4 Qwen3-Next GDN projections, HF weight mapping and explicit recurrent state. + +The module is shared by engine adapters. It does not install a vLLM backend or +claim model-level acceptance. The process group must contain exactly four ranks. +""" + +from collections.abc import Mapping, Sequence + +import torch +import torch.distributed as dist +from torch import nn + +from rl_engine.integrations.engines.train.vllm.qwen3_next_provider import ( + GDNProviderConfig, + GDNState, + shared_gdn_core, +) +from rl_engine.models.qwen3_next.qwen3_next_tp import _CopyToTP, _rank, _ReduceFromTP +from rl_engine.validation.common.tensor_identity import tensor_bitwise_equal + +_GLOBAL_SHAPES = { + "in_proj_qkvz.weight": (12288, 2048), + "in_proj_ba.weight": (64, 2048), + "conv1d.weight": (8192, 1, 4), + "A_log": (32,), + "dt_bias": (32,), + "norm.weight": (128,), + "out_proj.weight": (2048, 4096), +} +_LOCAL_SHAPES = { + "in_proj_qkvz.weight": (3072, 2048), + "in_proj_ba.weight": (16, 2048), + "conv1d.weight": (2048, 1, 4), + "A_log": (8,), + "dt_bias": (8,), + "norm.weight": (128,), + "out_proj.weight": (2048, 1024), +} + + +def _validate_weights(weights, shapes): + if set(weights) != set(shapes): + raise ValueError("GDN weight names do not match the explicit HF contract") + device = weights["A_log"].device + for name, shape in shapes.items(): + value = weights[name] + if value.shape != shape or value.dtype != torch.bfloat16 or value.device != device: + raise ValueError(f"{name} must be BF16 {shape} on {device}") + + +def shard_gdn_weights(weights: Mapping[str, torch.Tensor], rank: int) -> dict[str, torch.Tensor]: + """Shard official HF GQA-interleaved projections and Q/K/V-segmented conv.""" + _rank(rank) + _validate_weights(weights, _GLOBAL_SHAPES) + result = {} + for name, value in weights.items(): + if name == "norm.weight": + shard = value + elif name == "out_proj.weight": + shard = value[:, rank * 1024 : (rank + 1) * 1024] + elif name == "conv1d.weight": + q, k, v = value.split((2048, 2048, 4096), dim=0) + shard = torch.cat([part.chunk(4, dim=0)[rank] for part in (q, k, v)]) + else: + shard = value.chunk(4, dim=0)[rank] + result[name] = shard.contiguous() + return result + + +def assemble_gdn_weights(shards: Sequence[Mapping[str, torch.Tensor]]) -> dict[str, torch.Tensor]: + """Inverse mapping; reject divergent replicated norm weights.""" + if len(shards) != 4: + raise ValueError("Exactly four ordered TP shards are required") + for shard in shards: + _validate_weights(shard, _LOCAL_SHAPES) + result = {} + for name in _GLOBAL_SHAPES: + values = [shard[name] for shard in shards] + if name == "norm.weight": + if any(not tensor_bitwise_equal(value, values[0]) for value in values[1:]): + raise ValueError("Replicated GDN norm weights differ across TP ranks") + result[name] = values[0] + elif name == "conv1d.weight": + segments = [value.split((512, 512, 1024), dim=0) for value in values] + result[name] = torch.cat( + [segments[rank][part] for part in range(3) for rank in range(4)] + ) + else: + result[name] = torch.cat(values, dim=1 if name == "out_proj.weight" else 0) + return result + + +class TP4GDN(nn.Module): + """Complete GDN projection/core/output block with differentiable TP semantics.""" + + def __init__(self, *, group, device, config: GDNProviderConfig, eps: float = 1e-6): + super().__init__() + if not dist.is_initialized() or dist.get_world_size(group) != 4: + raise ValueError("TP4GDN requires an initialized four-rank process group") + config.validate() + self.group = dist.group.WORLD if group is None else group + self.provider_config, self.eps = config, eps + self.rank = dist.get_rank(group) + factory = {"device": device, "dtype": torch.bfloat16} + self.in_proj_qkvz = nn.Linear(2048, 3072, bias=False, **factory) + self.in_proj_ba = nn.Linear(2048, 16, bias=False, **factory) + self.conv1d = nn.Conv1d(2048, 2048, 4, groups=2048, bias=False, **factory) + self.A_log = nn.Parameter(torch.zeros(8, **factory)) + self.dt_bias = nn.Parameter(torch.zeros(8, **factory)) + self.norm = nn.Module() + self.norm.register_parameter("weight", nn.Parameter(torch.ones(128, **factory))) + self.out_proj = nn.Linear(1024, 2048, bias=False, **factory) + + def load_hf_weights(self, weights: Mapping[str, torch.Tensor]): + self.load_state_dict(shard_gdn_weights(weights, self.rank), strict=True) + + def initial_state(self, blocks: int) -> GDNState: + if blocks < 2: + raise ValueError("State must include a reserved slot and at least one active slot") + return GDNState( + self.A_log.new_zeros((blocks, 2048, 3)), + torch.zeros(blocks, 8, 128, 128, dtype=torch.float32, device=self.A_log.device), + ) + + def forward(self, hidden, state, indices, cu_seqlens): + from rl_engine.models.qwen3_next.qwen3_next_forward import shared_linear + + if hidden.ndim != 2 or hidden.shape[1] != 2048 or hidden.dtype != torch.bfloat16: + raise ValueError("GDN hidden states must be BF16 [tokens, 2048]") + copied = _CopyToTP.apply(hidden, self.group) + qkvz = shared_linear(copied, self.in_proj_qkvz.weight).reshape(-1, 4, 768) + ba = shared_linear(copied, self.in_proj_ba.weight).reshape(-1, 4, 4) + q, k, v, z = qkvz.split((128, 128, 256, 256), dim=-1) + b, a = ba.split(2, dim=-1) + tokens = hidden.shape[0] + packed = torch.cat( + [value.reshape(tokens, width) for value, width in ((q, 512), (k, 512), (v, 1024))], + dim=-1, + ) + # The norm weight is shared by heads on all four ranks; its gradient + # sums contributions from all local head partitions exactly once. + norm_weight = _CopyToTP.apply(self.norm.weight, self.group) + output, next_state = shared_gdn_core( + packed, + a.reshape(tokens, 8), + b.reshape(tokens, 8), + z.reshape(tokens, 8, 128), + self.A_log.float(), + self.dt_bias.float(), + self.conv1d.weight[:, 0], + norm_weight, + state, + indices, + cu_seqlens, + config=self.provider_config, + num_k_heads=4, + eps=self.eps, + ) + local = shared_linear(output.reshape(tokens, 1024), self.out_proj.weight) + return _ReduceFromTP.apply(local, self.group), next_state 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/common/tensor_identity.py b/rl_engine/validation/common/tensor_identity.py new file mode 100644 index 000000000..5d05c0fe9 --- /dev/null +++ b/rl_engine/validation/common/tensor_identity.py @@ -0,0 +1,38 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Finite tensor identity in the original storage dtype, including signed zero.""" + +import torch + + +def tensor_bitwise_equal(actual: torch.Tensor, expected: torch.Tensor) -> bool: + """Compare logical values by bytes without widening or accepting NaN/Inf. + + Strides need not match. A cross-device comparison copies expected bytes to + the actual tensor's device; it never converts either tensor's dtype. + """ + if actual.shape != expected.shape or actual.dtype != expected.dtype: + return False + for value in (actual, expected): + if (value.is_floating_point() or value.is_complex()) and not bool( + torch.isfinite(value).all() + ): + return False + actual_bytes = actual.detach().contiguous().reshape(-1).view(torch.uint8) + expected_bytes = expected.detach().contiguous().reshape(-1).view(torch.uint8) + if actual_bytes.device != expected_bytes.device: + expected_bytes = expected_bytes.to(actual_bytes.device) + return bool(torch.equal(actual_bytes, expected_bytes)) + + +def assert_tensor_bitwise_equal( + actual: torch.Tensor, expected: torch.Tensor, *, name: str = "tensor" +) -> None: + """Require finite, same-dtype, same-shape raw-bit identity.""" + if not tensor_bitwise_equal(actual, expected): + raise AssertionError( + f"{name}: finite raw-bit identity failed; " + f"actual={tuple(actual.shape)}/{actual.dtype}, " + f"expected={tuple(expected.shape)}/{expected.dtype}" + ) diff --git a/rl_engine/validation/models/qwen3_next_attention_prior_art.py b/rl_engine/validation/models/qwen3_next_attention_prior_art.py new file mode 100644 index 000000000..db705b56f --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_attention_prior_art.py @@ -0,0 +1,601 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Prior-art comparison for the Qwen3-Next TP4 full-attention core (D=256, GQA 4:1). + +Every candidate computes causal attention for one rank of Qwen3-Next's TP4 +layout: 4 query heads over 1 KV head, head dim 256, scale 1/16, queries aligned +to the end of the KV sequence. For each candidate the report records + +* **batch invariance** - a target sequence's output computed alone, and inside + a batch with three other sequences (placed first and last), must be bitwise + equal; for candidates with a backward, the same for its dq/dk/dv, in one + packed launch where the engine has a packed backward; +* **prefill/decode invariance** - the last 64 query rows, and the last single + row, computed against the full KV must equal the same rows of the full + prefill. This is the replay-vs-rollout boundary of RFC #428; +* **repeatability** - two identical calls are bitwise equal; +* **accuracy** - error against an FP64 evaluation; +* **performance** - median CUDA-event latency of prefill, decode and backward. + +A candidate that cannot be imported or launched in the pinned environment is +recorded as ``unavailable`` with the reason instead of failing the run. +""" + +from __future__ import annotations + +import statistics +from typing import Any, Callable + +import torch +import torch.nn.functional as F + +HQ, HKV, D = 4, 1, 256 +SCALE = 1.0 / 16 +TARGET = 1000 +COMPANIONS = (777, 1500, 64) +CHUNK = 64 +PREFILL_TOKENS = (512, 2048, 8192) +DECODE_KV = 8192 +BACKWARD_TOKENS = 2048 + + +# --------------------------------------------------------------------------- # +# Measurement helpers +# --------------------------------------------------------------------------- # + + +def _sample_us(fn: Callable[[], Any]) -> float: + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) * 1000.0 + + +def time_interleaved(calls: dict[str, Callable[..., Any]], warmup: int = 5, iters: int = 30): + keys = list(calls) + orders = (keys, list(reversed(keys))) + for i in range(warmup): + for key in orders[i % 2]: + calls[key]() + samples: dict[str, list[float]] = {key: [] for key in keys} + for i in range(iters): + for key in orders[i % 2]: + samples[key].append(_sample_us(calls[key])) + return {key: statistics.median(values) for key, values in samples.items()} + + +def _bitwise(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(a, b) + + +def _guard(key: str, factory: Callable[[], dict[str, Any]]) -> dict[str, Any]: + try: + return {"key": key, **factory()} + except Exception as exc: # noqa: BLE001 - record and continue + return {"key": key, "unavailable": f"{type(exc).__name__}: {str(exc)[:300]}"} + + +def make_sequence(length: int, seed: int, *, kv_length: int | None = None): + """q [Lq, HQ, D], k/v [Lk, HKV, D] in BF16 (token-major, as engines store them).""" + kv_length = length if kv_length is None else kv_length + g = torch.Generator(device="cuda").manual_seed(seed) + + def randn(*shape): + return torch.randn(shape, device="cuda", generator=g).to(torch.bfloat16) + + return randn(length, HQ, D), randn(kv_length, HKV, D), randn(kv_length, HKV, D) + + +def fp64_reference(q, k, v): + lq, lk = q.shape[0], k.shape[0] + qd, kd, vd = (t.double().transpose(0, 1) for t in (q, k, v)) + kd, vd = kd.expand(HQ, -1, -1), vd.expand(HQ, -1, -1) + scores = qd @ kd.transpose(1, 2) * SCALE + rows = torch.arange(lq, device=q.device)[:, None] + (lk - lq) + scores = scores.masked_fill(torch.arange(lk, device=q.device)[None, :] > rows, float("-inf")) + return (scores.softmax(-1) @ vd).transpose(0, 1) + + +# --------------------------------------------------------------------------- # +# Candidates +# +# ``run(seqs) -> outputs``: ``seqs`` is a list of (q, k, v) with queries aligned +# to the end of each KV sequence; every candidate computes the whole list in the +# way its engine batches it (one varlen launch, or one launch per sequence). +# ``grad(q, k, v, dy) -> (dq, dk, dv)`` when the candidate has a backward. +# --------------------------------------------------------------------------- # + + +def _rl_kernel(): + from rl_engine.models.qwen3_next.qwen3_next_forward import shared_attention + + def one(q, k, v): + out = shared_attention(*(t.transpose(0, 1).unsqueeze(0) for t in (q, k, v))) + return out.squeeze(0).transpose(0, 1) + + def run(seqs): + with torch.no_grad(): + return [one(*s) for s in seqs] + + def grad(q, k, v, dy): + leaves = [t.detach().clone().requires_grad_(True) for t in (q, k, v)] + return torch.autograd.grad(one(*leaves), leaves, dy) + + return { + "name": "RL-Kernel deterministic attention (one launch per sequence)", + "source": "rl_engine.models.qwen3_next.qwen3_next_forward.shared_attention", + "run": run, + "grad": grad, + } + + +def _sdpa(): + def one(q, k, v): + lq, lk = q.shape[0], k.shape[0] + qh, kh, vh = (t.transpose(0, 1).unsqueeze(0) for t in (q, k, v)) + if lq == lk: + out = F.scaled_dot_product_attention( + qh, kh, vh, is_causal=True, scale=SCALE, enable_gqa=True + ) + else: + rows = torch.arange(lq, device=q.device)[:, None] + (lk - lq) + mask = torch.arange(lk, device=q.device)[None, :] <= rows + out = F.scaled_dot_product_attention( + qh, kh, vh, attn_mask=mask, scale=SCALE, enable_gqa=True + ) + return out.squeeze(0).transpose(0, 1) + + def run(seqs): + with torch.no_grad(): + return [one(*s) for s in seqs] + + def grad(q, k, v, dy): + leaves = [t.detach().clone().requires_grad_(True) for t in (q, k, v)] + return torch.autograd.grad(one(*leaves), leaves, dy) + + return { + "name": "torch SDPA (default backend, one launch per sequence)", + "source": f"torch {torch.__version__} scaled_dot_product_attention", + "run": run, + "grad": grad, + } + + +def _cu(lengths): + out = [0] + for n in lengths: + out.append(out[-1] + n) + return torch.tensor(out, device="cuda", dtype=torch.int32) + + +def _split(out, lengths): + return list(out.split(list(lengths))) + + +def _vllm_fa(num_splits: int): + import vllm + from vllm.vllm_flash_attn import flash_attn_varlen_func + + def run(seqs): + lq = [s[0].shape[0] for s in seqs] + lk = [s[1].shape[0] for s in seqs] + q, k, v = (torch.cat([s[i] for s in seqs]) for i in range(3)) + out = flash_attn_varlen_func( + q, + k, + v, + max_seqlen_q=max(lq), + cu_seqlens_q=_cu(lq), + max_seqlen_k=max(lk), + cu_seqlens_k=_cu(lk), + softmax_scale=SCALE, + causal=True, + num_splits=num_splits, + fa_version=2, + ) + return _split(out, lq) + + return { + "name": f"vLLM FlashAttention-2 varlen, num_splits={num_splits}", + "source": f"vllm {vllm.__version__} vllm_flash_attn.flash_attn_varlen_func", + "run": run, + } + + +def _vllm_triton(): + import vllm + from vllm.v1.attention.ops.triton_unified_attention import unified_attention + + block = 16 + + def run(seqs): + lq = [s[0].shape[0] for s in seqs] + lk = [s[1].shape[0] for s in seqs] + blocks = [(n + block - 1) // block for n in lk] + k_cache = torch.zeros(sum(blocks), block, HKV, D, device="cuda", dtype=torch.bfloat16) + v_cache = torch.zeros_like(k_cache) + table = torch.zeros(len(seqs), max(blocks), device="cuda", dtype=torch.int32) + first = 0 + for i, (s, n) in enumerate(zip(seqs, blocks)): + pad = n * block - s[1].shape[0] + k_cache[first : first + n] = F.pad(s[1], (0, 0, 0, 0, 0, pad)).view(n, block, HKV, D) + v_cache[first : first + n] = F.pad(s[2], (0, 0, 0, 0, 0, pad)).view(n, block, HKV, D) + table[i, :n] = torch.arange(first, first + n, device="cuda") + first += n + q = torch.cat([s[0] for s in seqs]) + out = torch.empty_like(q) + unified_attention( + q, + k_cache, + v_cache, + out, + cu_seqlens_q=_cu(lq), + max_seqlen_q=max(lq), + seqused_k=torch.tensor(lk, device="cuda", dtype=torch.int32), + max_seqlen_k=max(lk), + softmax_scale=SCALE, + causal=True, + window_size=(-1, -1), + block_table=table, + softcap=0, + q_descale=None, + k_descale=None, + v_descale=None, + ) + return _split(out, lq) + + return { + # Without split-softmax buffers the 2D kernel runs for every batch, which + # is also the only kernel vLLM's batch-invariant mode allows. + "name": "vLLM Triton unified attention (2D kernel)", + "source": f"vllm {vllm.__version__} v1/attention/ops/triton_unified_attention", + "run": run, + } + + +def _flashinfer(): + import flashinfer + + workspace = torch.empty(256 << 20, device="cuda", dtype=torch.uint8) + wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(workspace, "NHD") + + def run(seqs): + lq = [s[0].shape[0] for s in seqs] + lk = [s[1].shape[0] for s in seqs] + wrapper.plan( + _cu(lq), + _cu(lk), + HQ, + HKV, + D, + causal=True, + sm_scale=SCALE, + q_data_type=torch.bfloat16, + ) + q, k, v = (torch.cat([s[i] for s in seqs]) for i in range(3)) + return _split(wrapper.run(q, k, v), lq) + + return { + "name": "FlashInfer BatchPrefillWithRaggedKVCache", + "source": f"flashinfer {flashinfer.__version__}", + "run": run, + } + + +def _bottom_right_mask(lq, lk, device): + """True where query row i (aligned to the end of the KV sequence) may not attend.""" + rows = torch.arange(lq, device=device)[:, None] + (lk - lq) + return torch.arange(lk, device=device)[None, :] > rows + + +def _fa4(): + """FlashAttention-4 (CuTe DSL), as vendored by vLLM 0.30.0: varlen forward and backward.""" + import vllm + from vllm.vllm_flash_attn.cute import interface as fa4 + + def call(q, k, v, lq, lk): + out = fa4.flash_attn_varlen_func( + q, + k, + v, + cu_seqlens_q=_cu(lq), + cu_seqlens_k=_cu(lk), + max_seqlen_q=max(lq), + max_seqlen_k=max(lk), + softmax_scale=SCALE, + causal=True, + ) + return out[0] if isinstance(out, tuple) else out + + def run(seqs): + lq, lk = [s[0].shape[0] for s in seqs], [s[1].shape[0] for s in seqs] + q, k, v = (torch.cat([s[i] for s in seqs]) for i in range(3)) + with torch.no_grad(): + return _split(call(q, k, v, lq, lk), lq) + + def grad_batch(seqs, dys): + lq, lk = [s[0].shape[0] for s in seqs], [s[1].shape[0] for s in seqs] + leaves = [torch.cat([s[i] for s in seqs]).detach().requires_grad_(True) for i in range(3)] + dq, dk, dv = torch.autograd.grad(call(*leaves, lq, lk), leaves, torch.cat(dys)) + return list(zip(_split(dq, lq), _split(dk, lk), _split(dv, lk))) + + return { + "name": "FlashAttention-4 (CuTe DSL) varlen, vendored by vLLM", + "source": f"vllm {vllm.__version__} vllm_flash_attn/cute " + "(its head-dim-256 SM100 backward has no deterministic mode)", + "run": run, + "grad_batch": grad_batch, + } + + +def _te(training: bool): + """Transformer Engine DotProductAttention on THD-packed sequences. + + With cuDNN 9.20 on SM100, TE has a fused (cuDNN) head-dim-256 backend only for + inference; in training mode it falls back to its unfused PyTorch path. Both are + measured, the fused one forward-only. + """ + import transformer_engine + import transformer_engine.pytorch as te + + attention = te.DotProductAttention( + HQ, + D, + num_gqa_groups=HKV, + attention_dropout=0.0, + qkv_format="thd", + attn_mask_type="padding_causal_bottom_right", + softmax_scale=SCALE, + ) + attention.train(training) + + def call(q, k, v, lq, lk): + out = attention( + q, + k, + v, + cu_seqlens_q=_cu(lq), + cu_seqlens_kv=_cu(lk), + max_seqlen_q=max(lq), + max_seqlen_kv=max(lk), + ) + return out.view(q.shape[0], HQ, D) + + def run(seqs): + lq, lk = [s[0].shape[0] for s in seqs], [s[1].shape[0] for s in seqs] + q, k, v = (torch.cat([s[i] for s in seqs]) for i in range(3)) + with torch.no_grad(): + return _split(call(q, k, v, lq, lk), lq) + + def grad_batch(seqs, dys): + lq, lk = [s[0].shape[0] for s in seqs], [s[1].shape[0] for s in seqs] + leaves = [torch.cat([s[i] for s in seqs]).detach().requires_grad_(True) for i in range(3)] + dq, dk, dv = torch.autograd.grad(call(*leaves, lq, lk), leaves, torch.cat(dys)) + return list(zip(_split(dq, lq), _split(dk, lk), _split(dv, lk))) + + version = f"transformer-engine {transformer_engine.__version__}" + if not training: + return { + "name": "Transformer Engine DotProductAttention (THD), inference: cuDNN fused", + "source": f"{version} (no fused head-dim-256 backward on SM100)", + "run": run, + } + return { + "name": "Transformer Engine DotProductAttention (THD), training: unfused fallback", + "source": f"{version} (selected for head-dim-256 training on SM100)", + "run": run, + "grad_batch": grad_batch, + } + + +def _megatron_local(): + """Megatron-core's own (non-TE) DotProductAttention; it has no packed-sequence path.""" + import os + + import megatron.core + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.dot_product_attention import DotProductAttention + from megatron.core.transformer.enums import AttnMaskType + from megatron.core.transformer.transformer_config import TransformerConfig + + if not dist.is_initialized(): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29532") + device = torch.device("cuda", torch.cuda.current_device()) + dist.init_process_group("nccl", rank=0, world_size=1, device_id=device) + if not parallel_state.model_parallel_is_initialized(): + parallel_state.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(0) + config = TransformerConfig( + num_layers=1, + hidden_size=HQ * D, + num_attention_heads=HQ, + num_query_groups=HKV, + kv_channels=D, + attention_dropout=0.0, + bf16=True, + params_dtype=torch.bfloat16, + ) + attention = DotProductAttention( + config, + layer_number=1, + attn_mask_type=AttnMaskType.padding, + attention_type="self", + softmax_scale=SCALE, + ) + + def one(q, k, v): + mask = _bottom_right_mask(q.shape[0], k.shape[0], q.device)[None, None] + out = attention(q.unsqueeze(1), k.unsqueeze(1), v.unsqueeze(1), mask) + return out.view(q.shape[0], HQ, D) + + def run(seqs): + with torch.no_grad(): + return [one(*s) for s in seqs] + + def grad(q, k, v, dy): + leaves = [t.detach().clone().requires_grad_(True) for t in (q, k, v)] + return torch.autograd.grad(one(*leaves), leaves, dy) + + return { + "name": "Megatron-core DotProductAttention (local, one launch per sequence)", + "source": f"megatron-core {megatron.core.__version__}", + "run": run, + "grad": grad, + } + + +def candidates() -> list[dict[str, Any]]: + return [ + _guard("rl_kernel_cuda", _rl_kernel), + _guard("torch_sdpa", _sdpa), + _guard("vllm_fa2_auto", lambda: _vllm_fa(0)), + _guard("vllm_fa2_split1", lambda: _vllm_fa(1)), + _guard("vllm_triton_2d", _vllm_triton), + _guard("flashinfer", _flashinfer), + _guard("fa4_cute", _fa4), + _guard("te_fused", lambda: _te(False)), + _guard("te_training", lambda: _te(True)), + _guard("megatron_local", _megatron_local), + ] + + +# --------------------------------------------------------------------------- # +# Checks +# --------------------------------------------------------------------------- # + + +def _batch_bi(cand, target, others) -> dict[str, bool]: + alone = cand["run"]([target])[0] + first = cand["run"]([target, *others])[0] + last = cand["run"]([*others, target])[-1] + return {"first": _bitwise(first, alone), "last": _bitwise(last, alone)} + + +def _prefill_decode(cand, target) -> dict[str, bool]: + q, k, v = target + full = cand["run"]([target])[0] + result = {} + for rows in (CHUNK, 1): + part = cand["run"]([(q[-rows:], k, v)])[0] + result[f"last_{rows}"] = _bitwise(part, full[-rows:]) + return result + + +def has_backward(cand) -> bool: + return "grad" in cand or "grad_batch" in cand + + +def grad_batch(cand, seqs, dys): + """Each sequence's (dq, dk, dv): one packed launch where the engine has one.""" + if "grad_batch" in cand: + return cand["grad_batch"](seqs, dys) + return [cand["grad"](*s, dy) for s, dy in zip(seqs, dys)] + + +def _grad_bi(cand, target, others) -> dict[str, bool]: + """The target's dq/dk/dv alone vs. packed first and last with the other sequences.""" + g = torch.Generator(device="cuda").manual_seed(5) + + def dy_for(q): + return torch.randn(q.shape, device="cuda", generator=g).to(torch.bfloat16) + + dy = dy_for(target[0]) + dys = [dy_for(o[0]) for o in others] + alone = grad_batch(cand, [target], [dy])[0] + first = grad_batch(cand, [target, *others], [dy, *dys])[0] + last = grad_batch(cand, [*others, target], [*dys, dy])[-1] + return { + name: _bitwise(f, a) and _bitwise(l_, a) + for name, f, l_, a in zip(("dq", "dk", "dv"), first, last, alone) + } + + +def _one(cand, target, others) -> dict[str, Any]: + if "unavailable" in cand: + return cand + out = {"key": cand["key"], "name": cand["name"], "source": cand["source"]} + try: + out["batch_bitwise"] = _batch_bi(cand, target, others) + try: + out["prefill_decode_bitwise"] = _prefill_decode(cand, target) + except Exception as exc: # noqa: BLE001 - recorded; the other checks still run + out["prefill_decode_failed"] = f"{type(exc).__name__}: {str(exc)[:300]}" + out["prefill_decode_bitwise"] = {"last_64": False, "last_1": False} + out["repeatable"] = _bitwise(cand["run"]([target])[0], cand["run"]([target])[0]) + if has_backward(cand): + out["backward_batch_bitwise"] = _grad_bi(cand, target, others) + ref = fp64_reference(*target) + got = cand["run"]([target])[0].double() + out["accuracy"] = { + "rel_l2_vs_fp64": float((got - ref).norm() / ref.norm()), + "max_abs_vs_fp64": float((got - ref).abs().max()), + } + except Exception as exc: # noqa: BLE001 - a crash is a result, not a harness failure + out["failed"] = f"{type(exc).__name__}: {str(exc)[:300]}" + out["batch_invariant"] = bool( + "failed" not in out + and all(out["batch_bitwise"].values()) + and all(out["prefill_decode_bitwise"].values()) + and all(out.get("backward_batch_bitwise", {"": True}).values()) + ) + return out + + +def _latency(cands) -> dict[str, Any]: + ready = [c for c in cands if "unavailable" not in c] + + def calls_for(seqs): + calls = {} + for cand in ready: + try: + cand["run"](seqs) + except Exception: # noqa: BLE001 - recorded by the BI section + continue + calls[cand["key"]] = lambda c=cand: c["run"](seqs) + return calls + + prefill = { + str(n): time_interleaved(calls_for([make_sequence(n, 300 + n)])) for n in PREFILL_TOKENS + } + decode = time_interleaved(calls_for([make_sequence(1, 401, kv_length=DECODE_KV)])) + q, k, v = make_sequence(BACKWARD_TOKENS, 402) + dy = torch.randn_like(q) + backward = time_interleaved( + { + c["key"]: (lambda c=c: grad_batch(c, [(q, k, v)], [dy])) + for c in ready + if has_backward(c) + }, + warmup=3, + iters=15, + ) + return { + "prefill_us": prefill, + f"decode_us_kv{DECODE_KV}": decode, + f"forward_plus_backward_us_{BACKWARD_TOKENS}": backward, + } + + +def attention_report(only=None) -> dict[str, Any]: + """Every candidate, or only the keys in ``only`` (run each Python environment separately).""" + target = make_sequence(TARGET, 1) + others = [make_sequence(n, 10 + i) for i, n in enumerate(COMPANIONS)] + cands = [c for c in candidates() if only is None or c["key"] in only] + if only is not None and set(only) - {c["key"] for c in cands}: + raise ValueError(f"Unknown candidates: {sorted(set(only) - {c['key'] for c in cands})}") + return { + "op": "qwen3_next_tp4_attention", + "shape": {"q_heads": HQ, "kv_heads": HKV, "head_dim": D, "scale": SCALE}, + "target_tokens": TARGET, + "companion_tokens": list(COMPANIONS), + "candidates": [_one(c, target, others) for c in cands], + "latency": _latency(cands), + } diff --git a/rl_engine/validation/models/qwen3_next_moe_prior_art.py b/rl_engine/validation/models/qwen3_next_moe_prior_art.py new file mode 100644 index 000000000..7c0cc33c9 --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_moe_prior_art.py @@ -0,0 +1,688 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Prior-art comparison for the Qwen3-Next routed MoE (RFC #428 ``moe_route_combine_contract``). + +Every candidate computes the routed half of ``Qwen3NextSparseMoeBlock`` at the +official TP1 shape (H=2048, 512 experts, top-10, expert width 512): router +logits, top-10 selection with renormalised weights, the selected SwiGLU +experts and the weighted combine. The shared expert is left out because every +candidate evaluates it as the same dense MLP. + +For each candidate the report records + +* **routing batch invariance** - the expert ids and weights of a fixed block + of probe tokens computed alone, and embedded first and last in larger + batches, must be bitwise equal; +* **output batch invariance** - the same for the routed output; +* **dx / dW batch invariance** (candidates with a backward) - the probe rows' + input gradient must not change when unrelated rows join the batch, and the + expert weight gradient must not change when rows whose output gradient is + zero join it (mathematically neither can); +* **repeatability** - two identical calls are bitwise equal; +* **accuracy** - error against an FP64 evaluation of the HF formula with FP64 + routing, plus the fraction of tokens whose selected expert set differs; +* **performance** - median CUDA-event latency, candidates interleaved. + +A candidate that cannot be imported or launched in the pinned environment is +recorded as ``unavailable`` with the reason instead of failing the run. +""" + +from __future__ import annotations + +import os +import statistics +from contextlib import contextmanager +from typing import Any, Callable + +import torch +import torch.nn.functional as F + +HIDDEN = 2048 +EXPERTS = 512 +TOPK = 10 +WIDTH = 512 +PROBE = 8 +BI_SIZES = (16, 64, 256, 1024) +PERF_TOKENS = (1, 8, 64, 256, 1024, 4096) +BACKWARD_TOKENS = (64, 1024) + + +# --------------------------------------------------------------------------- # +# Measurement helpers +# --------------------------------------------------------------------------- # + + +def _sample_us(fn: Callable[[], Any]) -> float: + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) * 1000.0 + + +def time_interleaved(calls: dict[str, Callable[..., Any]], warmup: int = 5, iters: int = 30): + """Median latency per candidate; the order alternates so none always runs first.""" + keys = list(calls) + orders = (keys, list(reversed(keys))) + for i in range(warmup): + for key in orders[i % 2]: + calls[key]() + samples: dict[str, list[float]] = {key: [] for key in keys} + for i in range(iters): + for key in orders[i % 2]: + samples[key].append(_sample_us(calls[key])) + return {key: statistics.median(values) for key, values in samples.items()} + + +def _bitwise(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(a, b) + + +def _rel(value: torch.Tensor, reference: torch.Tensor) -> float: + reference = reference.double() + return float((value.double() - reference).norm() / reference.norm().clamp_min(1e-300)) + + +def _guard(key: str, factory: Callable[[], dict[str, Any]]) -> dict[str, Any]: + try: + return {"key": key, **factory()} + except Exception as exc: # noqa: BLE001 - record and continue + return {"key": key, "unavailable": f"{type(exc).__name__}: {str(exc)[:300]}"} + + +@contextmanager +def _env(name: str, value: str | None): + old = os.environ.get(name) + if value is None: + os.environ.pop(name, None) + else: + os.environ[name] = value + try: + yield + finally: + if old is None: + os.environ.pop(name, None) + else: + os.environ[name] = old + + +# --------------------------------------------------------------------------- # +# Inputs and the FP64 reference +# --------------------------------------------------------------------------- # + + +def make_weights(seed: int = 0) -> dict[str, torch.Tensor]: + g = torch.Generator(device="cuda").manual_seed(seed) + + def randn(*shape, scale): + return (torch.randn(shape, device="cuda", generator=g) * scale).to(torch.bfloat16) + + return { + "router": randn(EXPERTS, HIDDEN, scale=0.02), + "gate_up": randn(EXPERTS, 2 * WIDTH, HIDDEN, scale=0.02), + "down": randn(EXPERTS, HIDDEN, WIDTH, scale=0.02), + } + + +def make_tokens(count: int, seed: int) -> torch.Tensor: + g = torch.Generator(device="cuda").manual_seed(seed) + return torch.randn(count, HIDDEN, device="cuda", generator=g).to(torch.bfloat16) + + +def fp64_reference(x: torch.Tensor, w: dict[str, torch.Tensor]): + """HF formula in FP64 with FP64 routing; returns (output, expert ids).""" + x64 = x.double() + probs = (x64 @ w["router"].double().T).softmax(-1) + weights, ids = probs.topk(TOPK, dim=-1) + weights = weights / weights.sum(-1, keepdim=True) + out = torch.zeros_like(x64) + gate_up, down = w["gate_up"], w["down"] + for expert in ids.unique().tolist(): + token, slot = torch.where(ids == expert) + gate, up = (x64[token] @ gate_up[expert].double().T).chunk(2, dim=-1) + y = (F.silu(gate) * up) @ down[expert].double().T + out.index_add_(0, token, y * weights[token, slot, None]) + return out, ids + + +# --------------------------------------------------------------------------- # +# Candidates +# +# ``route(x) -> (ids [T,10], weights [T,10])`` exposes the routing decision; +# ``fwd(x) -> [T, H]`` is the routed output; ``graph(x) -> (y, leaves)`` (when +# the candidate is differentiable) returns the output and [x, gate_up, down]. +# --------------------------------------------------------------------------- # + + +def _rl_kernel(w): + from rl_engine.models.qwen3_next import qwen3_next_forward as provider + + def route(x): + routes = provider.stable_top10_routes(provider.shared_router(x, w["router"])) + return routes.indices, routes.weights + + def fwd(x): + with torch.no_grad(): + return provider.shared_moe(x, w["router"], w["gate_up"], w["down"])[0] + + gate_up = w["gate_up"].detach().clone().requires_grad_(True) + down = w["down"].detach().clone().requires_grad_(True) + + def graph(x): + leaf = x.detach().clone().requires_grad_(True) + y, _ = provider.shared_moe(leaf, w["router"], gate_up, down) + return y, [leaf, gate_up, down] + + return { + "name": "rl-kernel shared_moe", + "source": "rl_engine.models.qwen3_next.qwen3_next_forward.shared_moe", + "route": route, + "fwd": fwd, + "graph": graph, + } + + +def _hf_modules(w, device="cuda"): + from transformers import Qwen3NextConfig + from transformers.models.qwen3_next import modeling_qwen3_next as hf + + config = Qwen3NextConfig( + hidden_size=HIDDEN, + num_experts=EXPERTS, + num_experts_per_tok=TOPK, + moe_intermediate_size=WIDTH, + norm_topk_prob=True, + ) + router = hf.Qwen3NextTopKRouter(config).to(device=device, dtype=torch.bfloat16) + experts = hf.Qwen3NextExperts(config).to(device=device, dtype=torch.bfloat16) + with torch.no_grad(): + router.weight.copy_(w["router"]) + experts.gate_up_proj.copy_(w["gate_up"]) + experts.down_proj.copy_(w["down"]) + return router, experts + + +def _hf(w): + import transformers + + router, experts = _hf_modules(w) + + def route(x): + _, weights, ids = router(x) + return ids, weights + + def call(x, mod): + _, weights, ids = router(x) + return mod(x, ids, weights) + + def fwd(x): + with torch.no_grad(): + return call(x, experts) + + def graph(x): + leaf = x.detach().clone().requires_grad_(True) + return call(leaf, experts), [leaf, experts.gate_up_proj, experts.down_proj] + + return { + "name": "HF transformers Qwen3NextExperts (eager)", + "source": f"transformers {transformers.__version__} modeling_qwen3_next", + "route": route, + "fwd": fwd, + "graph": graph, + } + + +def _vllm(w, batch_invariant: bool): + import vllm + from vllm.model_executor.layers.fused_moe.fused_moe import fused_experts + from vllm.model_executor.layers.fused_moe.router.fused_topk_router import fused_topk + + flag = "1" if batch_invariant else "0" + linear = F.linear + if batch_invariant: + from vllm.model_executor.determinism.batch_invariant import linear_batch_invariant + + linear = linear_batch_invariant + + def route(x): + with _env("VLLM_BATCH_INVARIANT", flag): + # vLLM's Qwen3-Next gate is a BF16 ReplicatedLinear; batch-invariant + # mode replaces its aten GEMM with vLLM's persistent Triton matmul. + logits = linear(x, w["router"]) + weights, ids, _ = fused_topk(x, logits, TOPK, renormalize=True) + return ids.long(), weights + + def fwd(x): + with torch.no_grad(), _env("VLLM_BATCH_INVARIANT", flag): + ids, weights = route(x) + return fused_experts(x, w["gate_up"], w["down"], weights, ids.int()) + + return { + "name": f"vLLM fused_topk + fused_experts (Triton), VLLM_BATCH_INVARIANT={flag}", + "source": f"vllm {vllm.__version__} model_executor/layers/fused_moe", + "route": route, + "fwd": fwd, + } + + +def _flashinfer(w): + import flashinfer + from flashinfer.fused_moe import cutlass_fused_moe + + # CUTLASS SwiGLU takes fc1 as [up; gate] (the opposite half order of HF). + gate, up = w["gate_up"].chunk(2, dim=1) + fc1 = torch.cat((up, gate), dim=1).contiguous() + + def route(x): + probs = F.linear(x, w["router"]).float().softmax(-1) + weights, ids = probs.topk(TOPK, dim=-1) + return ids, weights / weights.sum(-1, keepdim=True) + + def fwd(x): + with torch.no_grad(): + ids, weights = route(x) + out = cutlass_fused_moe( + x, ids.int(), weights, fc1, w["down"], torch.bfloat16, quant_scales=[] + ) + return out[0] if isinstance(out, (list, tuple)) else out + + return { + "name": "FlashInfer cutlass_fused_moe (routes from torch.topk)", + "source": f"flashinfer {flashinfer.__version__} fused_moe.cutlass_fused_moe", + "route": route, + "fwd": fwd, + } + + +def _megatron_layer(w): + """Megatron-core MoELayer configured as VIME runs Qwen3-Next. + + VIME's ``scripts/models/qwen3-next-80B-A3B.sh``: + + Softmax router computed in FP32, top-10 of 512, all-to-all dispatcher, TE + grouped GEMM and TE fused permute, no auxiliary loss. Single process, TP1/EP1. + """ + import os + + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.spec_utils import build_module + from megatron.core.transformer.transformer_config import TransformerConfig + + if not dist.is_initialized(): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29531") + device = torch.device("cuda", torch.cuda.current_device()) + dist.init_process_group("nccl", rank=0, world_size=1, device_id=device) + if not parallel_state.model_parallel_is_initialized(): + parallel_state.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(0) + config = TransformerConfig( + num_layers=1, + hidden_size=HIDDEN, + num_attention_heads=16, + ffn_hidden_size=WIDTH, + moe_ffn_hidden_size=WIDTH, + num_moe_experts=EXPERTS, + moe_router_topk=TOPK, + moe_router_score_function="softmax", + moe_router_dtype="fp32", + moe_token_dispatcher_type="alltoall", + moe_grouped_gemm=True, + moe_permute_fusion=True, + moe_aux_loss_coeff=0.0, + moe_router_load_balancing_type="none", + add_bias_linear=False, + gated_linear_unit=True, + activation_func=F.silu, + bf16=True, + params_dtype=torch.bfloat16, + ) + spec = get_moe_module_spec(use_te=True, num_experts=EXPERTS, moe_grouped_gemm=True) + layer = build_module(spec, config=config).cuda() + with torch.no_grad(): + layer.router.weight.copy_(w["router"]) + for expert in range(EXPERTS): + getattr(layer.experts.linear_fc1, f"weight{expert}").copy_(w["gate_up"][expert]) + getattr(layer.experts.linear_fc2, f"weight{expert}").copy_(w["down"][expert]) + return layer + + +def _megatron(w): + import megatron.core + import transformer_engine + + layer = _megatron_layer(w) + gate_up = [getattr(layer.experts.linear_fc1, f"weight{e}") for e in range(EXPERTS)] + down = [getattr(layer.experts.linear_fc2, f"weight{e}") for e in range(EXPERTS)] + + def route(x): + with torch.no_grad(): + probs, routing_map = layer.router(x) + weights, ids = probs.topk(TOPK, dim=-1) + if not bool(routing_map.gather(1, ids).all()): + raise RuntimeError("Megatron routing map disagrees with its routed probabilities") + return ids, weights + + def call(x): + return layer(x.unsqueeze(1))[0].squeeze(1) + + def fwd(x): + with torch.no_grad(): + return call(x) + + def grads(x, dy): + leaf = x.detach().clone().requires_grad_(True) + dx, *dw = torch.autograd.grad(call(leaf), [leaf, *gate_up, *down], dy) + return dx, torch.stack(dw[:EXPERTS]), torch.stack(dw[EXPERTS:]) + + return { + "name": "Megatron-core MoELayer + TE grouped GEMM (VIME's Qwen3-Next config)", + "source": f"megatron-core {megatron.core.__version__}, " + f"transformer-engine {transformer_engine.__version__}", + "route": route, + "fwd": fwd, + "grads": grads, + } + + +def _sglang_shims(): + """Make SGLang's Triton MoE path importable next to torch 2.13. + + ``sgl_kernel`` 0.3.21 is built against another libtorch ABI and cannot load + here. On CUDA, SGLang's Triton MoE calls two of its kernels; both are + replaced by SGLang's own implementations of the same operation: the Triton + ``moe_sum_reduce_triton`` and the JIT ``moe_align_block_size`` that SGLang + registers with the same signature. Every other ``sgl_kernel`` symbol only + has to import; calling one raises. + """ + import importlib.abc + import importlib.machinery + import sys + import types + + class Missing: + def __init__(self, name): + self.name = name + + def __getattr__(self, name): + if name.startswith("__"): + raise AttributeError(name) + return Missing(f"{self.name}.{name}") + + def __call__(self, *args, **kwargs): + raise RuntimeError(f"{self.name} called, but sgl_kernel cannot load next to torch 2.13") + + class Stub(types.ModuleType): + def __getattr__(self, name): + if name.startswith("__"): + raise AttributeError(name) + return Missing(f"{self.__name__}.{name}") + + class Finder(importlib.abc.MetaPathFinder, importlib.abc.Loader): + def find_spec(self, name, path=None, target=None): + if name in ("sgl_kernel", "gguf") or name.startswith("sgl_kernel."): + return importlib.machinery.ModuleSpec(name, self, is_package=True) + + def create_module(self, spec): + module = Stub(spec.name) + module.__path__ = [] + return module + + def exec_module(self, module): + pass + + if not any(type(f).__name__ == "Finder" for f in sys.meta_path): + sys.meta_path.insert(0, Finder()) + import sgl_kernel + import sglang.kernels.ops.moe as moe_ops + from sglang.kernels.ops.moe.fused_moe_triton_kernels import moe_sum_reduce_triton + from sglang.kernels.spec import KernelBackend + + sgl_kernel.moe_sum_reduce = moe_sum_reduce_triton + lookup = moe_ops.get_kernel + moe_ops.get_kernel = lambda op, backend: lookup( + op, KernelBackend.JIT if op == "moe.moe_align_block_size" else backend + ) + + +def _sglang(w, deterministic: bool): + import sglang + + _sglang_shims() + from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig + from sglang.srt.layers.moe.moe_runner.triton_utils import fused_moe as fm + from sglang.srt.layers.moe.moe_runner.triton_utils import fused_moe_triton_config as fc + from sglang.srt.layers.moe.topk import StandardTopKOutput, fused_topk_torch_native + + class Bag: + def __init__(self, **values): + self.__dict__.update(values) + + def __getattr__(self, name): + return False + + # SGLang reads these switches from its published server config; publishing + # needs a full server (and sgl_kernel). Only the deterministic switch matters. + settings = Bag(deterministic=Bag(enable_deterministic_inference=deterministic), moe=Bag()) + from sglang.srt.batch_invariant_ops import batch_invariant_ops as bio + + def linear(x, weight): + # Deterministic mode: what enable_batch_invariant_mode() installs for aten::mm, + # called directly so the override does not leak into other candidates. + return bio.mm_batch_invariant(x, weight.t()) if deterministic else F.linear(x, weight) + + config = MoeRunnerConfig( + num_experts=EXPERTS, + num_local_experts=EXPERTS, + hidden_size=HIDDEN, + intermediate_size_per_partition=WIDTH, + top_k=TOPK, + params_dtype=torch.bfloat16, + inplace=False, + ) + + def set_mode(): + fm.get_exec = fc.get_exec = lambda: settings + fm.is_batch_invariant_mode_enabled = lambda: deterministic + # One process, no TP group: symmetric memory is disabled, the group unused. + fm.get_parallel = lambda: Bag(tp_group=None) + fm.is_allocation_symmetric = lambda: False + + def route(x): + set_mode() + logits = linear(x, w["router"]) + weights, ids = fused_topk_torch_native(x, logits, TOPK, renormalize=True)[:2] + return ids.long(), weights + + def fwd(x): + with torch.no_grad(): + ids, weights = route(x) + topk = StandardTopKOutput(weights, ids.int(), None) + return fm.fused_experts(x, w["gate_up"], w["down"], topk, config) + + mode = "deterministic inference" if deterministic else "default" + return { + "name": f"SGLang fused_moe (Triton), {mode}", + "source": f"sglang {sglang.__version__} moe_runner/triton_utils", + "route": route, + "fwd": fwd, + } + + +def candidates(w) -> list[dict[str, Any]]: + return [ + _guard("rl_kernel_cuda", lambda: _rl_kernel(w)), + _guard("hf_transformers", lambda: _hf(w)), + _guard("vllm_bi0", lambda: _vllm(w, False)), + _guard("vllm_bi1", lambda: _vllm(w, True)), + _guard("flashinfer_cutlass", lambda: _flashinfer(w)), + _guard("megatron_te", lambda: _megatron(w)), + _guard("sglang_triton", lambda: _sglang(w, False)), + _guard("sglang_deterministic", lambda: _sglang(w, True)), + ] + + +# --------------------------------------------------------------------------- # +# Checks +# --------------------------------------------------------------------------- # + + +def _embedded(x: torch.Tensor, size: int): + """Probe rows first in a batch of ``size``, and the same probe rows last.""" + first = x[:size] + last = x[:size].clone() + last[size - PROBE :] = x[:PROBE] + return first, last + + +def _rows_bi(fn, x) -> dict[str, bool]: + alone = fn(x[:PROBE]) + result = {} + for size in BI_SIZES: + first, last = _embedded(x, size) + a, b = fn(first), fn(last) + if isinstance(alone, tuple): + ok = all( + _bitwise(p[:PROBE], q) and _bitwise(r[size - PROBE :], q) + for p, r, q in zip(a, b, alone) + ) + else: + ok = _bitwise(a[:PROBE], alone) and _bitwise(b[size - PROBE :], alone) + result[str(size)] = bool(ok) + return result + + +def has_backward(cand) -> bool: + return "graph" in cand or "grads" in cand + + +def cand_grads(cand, rows, dy_rows): + """``(dx, d gate_up, d down)`` of one backward.""" + if "grads" in cand: + return cand["grads"](rows, dy_rows) + y, leaves = cand["graph"](rows) + return torch.autograd.grad(y, leaves, dy_rows) + + +def _grad_bi(cand, x) -> dict[str, Any]: + """dx of the probe rows, and dW with zero-gradient rows appended.""" + g = torch.Generator(device="cuda").manual_seed(7) + dy = torch.randn(x.shape, device="cuda", generator=g).to(torch.bfloat16) + + def grads(rows, dy_rows): + return cand_grads(cand, rows, dy_rows) + + alone = grads(x[:PROBE], dy[:PROBE]) + dx, dw = {}, {} + for size in BI_SIZES: + full = grads(x[:size], dy[:size]) + dx[str(size)] = bool(_bitwise(full[0][:PROBE], alone[0])) + padded = dy[:size].clone() + padded[PROBE:] = 0 + masked = grads(x[:size], padded) + dw[str(size)] = bool(all(_bitwise(m, a) for m, a in zip(masked[1:], alone[1:]))) + return {"dx_rows_bitwise": dx, "dweight_zero_rows_bitwise": dw} + + +def _accuracy(cand, x, w) -> dict[str, Any]: + ref, ref_ids = fp64_reference(x, w) + ids, _ = cand["route"](x) + same = ids.long().sort(-1).values == ref_ids.sort(-1).values + out = cand["fwd"](x) + return { + "rel_l2_vs_fp64": _rel(out, ref), + "max_abs_vs_fp64": float((out.double() - ref).abs().max()), + "tokens_with_different_expert_set": int((~same.all(-1)).sum()), + "tokens": int(x.shape[0]), + } + + +def _one(cand, x, w) -> dict[str, Any]: + if "unavailable" in cand: + return cand + out = {"key": cand["key"], "name": cand["name"], "source": cand["source"]} + try: + out["route_rows_bitwise"] = _rows_bi(cand["route"], x) + out["output_rows_bitwise"] = _rows_bi(cand["fwd"], x) + out["repeatable"] = bool(_bitwise(cand["fwd"](x[:256]), cand["fwd"](x[:256]))) + if has_backward(cand): + out.update(_grad_bi(cand, x)) + out["accuracy"] = _accuracy(cand, x[:256], w) + except Exception as exc: # noqa: BLE001 - a crash is a result, not a harness failure + out["failed"] = f"{type(exc).__name__}: {str(exc)[:300]}" + out["batch_invariant"] = bool( + "failed" not in out + and all(out["route_rows_bitwise"].values()) + and all(out["output_rows_bitwise"].values()) + and all(out.get("dx_rows_bitwise", {"": True}).values()) + and all(out.get("dweight_zero_rows_bitwise", {"": True}).values()) + ) + return out + + +def _latency(cands, tokens) -> dict[str, Any]: + ready = [c for c in cands if "unavailable" not in c] + forward = {} + for count in tokens: + x = make_tokens(count, seed=100 + count) + calls = {} + for cand in ready: + try: + cand["fwd"](x) + calls[cand["key"]] = lambda c=cand: c["fwd"](x) + except Exception: # noqa: BLE001 - recorded by the BI section + continue + forward[str(count)] = time_interleaved(calls) + backward = {} + for count in BACKWARD_TOKENS: + x = make_tokens(count, seed=200 + count) + dy = torch.randn_like(x) + calls = {} + for cand in ready: + if not has_backward(cand): + continue + + def step(c=cand): + cand_grads(c, x, dy) + + calls[cand["key"]] = step + backward[str(count)] = time_interleaved(calls, warmup=2, iters=10) + return {"forward_us": forward, "forward_plus_backward_us": backward} + + +def moe_report(only=None) -> dict[str, Any]: + """The report for every candidate, or only for the keys in ``only``. + + Some candidates need another Python environment (SGLang's pinned torch ABI, + Megatron's training stack); run each environment with ``only`` and merge. + """ + torch.manual_seed(0) + w = make_weights() + x = make_tokens(max(BI_SIZES), seed=1) + keys = [ + "rl_kernel_cuda", + "hf_transformers", + "vllm_bi0", + "vllm_bi1", + "flashinfer_cutlass", + "megatron_te", + "sglang_triton", + "sglang_deterministic", + ] + if only is not None and set(only) - set(keys): + raise ValueError(f"Unknown candidates: {sorted(set(only) - set(keys))}") + cands = [c for c in candidates(w) if only is None or c["key"] in only] + return { + "op": "qwen3_next_routed_moe", + "shape": {"hidden": HIDDEN, "experts": EXPERTS, "top_k": TOPK, "expert_width": WIDTH}, + "probe_rows": PROBE, + "candidates": [_one(c, x, w) for c in cands], + "latency": _latency(cands, PERF_TOKENS), + } 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_attention.py b/tests/models/qwen3_next/check_qwen3_next_attention.py new file mode 100644 index 000000000..9dfe6c267 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_attention.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Required Qwen3-Next CUDA attention gates; unavailable CUDA is a failure.""" + +import pytest +import torch + +from rl_engine.backends.cuda.attention.deterministic_attn import DeterministicAttentionOp +from rl_engine.backends.extension import _C +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def inputs(tokens, *, batch=1, query_tokens=None, device="cuda", grad=False): + assert torch.cuda.is_available(), "Mandatory attention CUDA gate cannot skip" + gen = torch.Generator(device=device).manual_seed(920) + query_tokens = tokens if query_tokens is None else query_tokens + return tuple( + torch.randn(batch, heads, length, 256, device=device, dtype=torch.bfloat16, generator=gen) + .mul_(0.25) + .requires_grad_(grad) + for heads, length in ((4, query_tokens), (1, tokens), (1, tokens)) + ) + + +@pytest.mark.parametrize("tokens", [8, 64, 256, 1024]) +@torch.no_grad() +def test_d256_full_chunk_decode_and_batch_are_raw_bit_equal(tokens): + q, k, v = inputs(tokens, batch=4) + op = DeterministicAttentionOp() + whole, lse = op.forward_with_lse(q, k, v) + order = torch.tensor([3, 1, 0, 2], device=q.device) + exact(op(q[order], k[order], v[order]), whole[order]) + for row in range(4): + exact(op(q[row : row + 1], k[row : row + 1], v[row : row + 1]), whole[row : row + 1]) + starts = [0, 1, 7, tokens - 1, tokens] + starts = sorted(set(starts)) + for start, end in zip(starts[:-1], starts[1:]): + out, part_lse = op.forward_with_lse(q[:, :, start:end], k[:, :, :end], v[:, :, :end]) + exact(out, whole[:, :, start:end]) + exact(part_lse, lse[:, :, start:end]) + + +def test_d256_backward_with_cached_prefix_matches_float64_formula(): + q, k, v = inputs(19, query_tokens=3, grad=True) + op = DeterministicAttentionOp() + out = op(q, k, v) + grad = torch.linspace(-1, 1, out.numel(), device=q.device).reshape_as(out).to(out.dtype) + out.backward(grad) + qr, kr, vr = (x.detach().double().requires_grad_() for x in (q, k, v)) + scores = qr @ kr.repeat_interleave(4, dim=1).transpose(-1, -2) / 16 + causal = ( + torch.arange(19, device=q.device)[None, :] <= torch.arange(16, 19, device=q.device)[:, None] + ) + expected = scores.masked_fill(~causal, -torch.inf).softmax(-1) @ vr.repeat_interleave(4, dim=1) + expected.backward(grad.double()) + # Formula accuracy is separate from the raw-bit provider comparisons above. + for actual, reference in zip((q, k, v), (qr, kr, vr)): + assert torch.isfinite(actual.grad).all() and actual.grad.abs().max() > 0 + torch.testing.assert_close( + actual.grad.float(), reference.grad.float(), atol=0.002, rtol=0.016 + ) + assert k.grad[:, :, :16].abs().max() > 0 + assert v.grad[:, :, :16].abs().max() > 0 + + +def test_native_backward_rejects_malformed_probability_and_gradient_buffers(): + q, k, v = inputs(8) + out, _, probabilities = _C.deterministic_attention_forward(q, k, v, True, 1 / 16, None) + with pytest.raises(RuntimeError, match="P must be FP32"): + _C.deterministic_attention_backward( + out, q, k, v, probabilities.to(q.dtype), True, 1 / 16, None + ) + with pytest.raises(RuntimeError, match="grad_output must have q shape"): + _C.deterministic_attention_backward( + out[:, :, :1], q, k, v, probabilities, True, 1 / 16, None + ) + with pytest.raises(RuntimeError, match="P must have shape"): + _C.deterministic_attention_backward( + out, q, k, v, probabilities[:, :, :1], True, 1 / 16, None + ) + + +def test_d256_input_device_guard_and_nondefault_stream(): + assert torch.cuda.device_count() >= 2, "Mandatory cross-device gate requires allocated GPUs" + previous = torch.cuda.current_device() + try: + q, k, v = inputs(8, device="cuda:1", grad=True) + op = DeterministicAttentionOp() + torch.cuda.set_device(1) + expected = op(q, k, v) + expected.sum().backward() + gradients = [x.grad.clone() for x in (q, k, v)] + for x in (q, k, v): + x.grad = None + torch.cuda.synchronize(1) + stream = torch.cuda.Stream(device=1) + with torch.cuda.stream(stream): + torch.cuda.set_device(0) + actual = op(q, k, v) + actual.sum().backward() + assert torch.cuda.current_device() == 0 + stream.synchronize() + exact(actual, expected) + for x, expected_grad in zip((q, k, v), gradients): + exact(x.grad, expected_grad) + with pytest.raises(ValueError, match="same device"): + op(q, k.to("cuda:0"), v) + finally: + torch.cuda.set_device(previous) diff --git a/tests/models/qwen3_next/check_qwen3_next_conv_bridge.py b/tests/models/qwen3_next/check_qwen3_next_conv_bridge.py new file mode 100644 index 000000000..963d8b630 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_conv_bridge.py @@ -0,0 +1,99 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The convolution provider must agree across batch, chunk and cache layouts.""" + +import pytest +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_conv import ( + _provider, + causal_conv_sequence, +) +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def inputs(lengths): + assert torch.cuda.is_available(), "Required provider tests cannot skip CUDA" + gen = torch.Generator(device="cuda").manual_seed(17) + + def rand(*shape): + return torch.randn(*shape, device="cuda", dtype=torch.bfloat16, generator=gen) + + return ( + rand(sum(lengths), 2048), + rand(len(lengths) + 1, 2048, 3), + rand(2048, 4), + torch.arange(1, len(lengths) + 1, device="cuda", dtype=torch.int32), + torch.tensor([0, *torch.tensor(lengths).cumsum(0).tolist()]), + ) + + +@torch.no_grad() +@pytest.mark.parametrize("lengths", [(1,), (64, 8, 0, 16), (1024, 256, 64, 8)]) +def test_varlen_matches_repeated_decode_and_preserves_inputs(lengths): + x, state, weight, indices, cu = inputs(lengths) + saved = [tensor.clone() for tensor in (x, state, weight)] + output, final_state = causal_conv_sequence(x, state, weight, indices, cu) + reference_state = state.clone() + reference_rows = [] + for seq in range(len(lengths)): + for token in range(cu[seq], cu[seq + 1]): + value = x[token : token + 1].clone() + reference_rows.append( + _provider()( + value, + reference_state, + weight, + activation="silu", + conv_state_indices=indices[seq : seq + 1], + ) + ) + exact(output, torch.cat(reference_rows)) + exact(final_state, reference_state) + for original, before in zip((x, state, weight), saved): + exact(original, before) + + +@torch.no_grad() +def test_chunked_and_reordered_sequences_match(): + x, state, weight, indices, cu = inputs((1024, 64)) + whole, whole_state = causal_conv_sequence(x, state, weight, indices, cu) + # Continue both sequences after 32 tokens, in the opposite batch order. + prompt = torch.cat((x[:32], x[1024:1056])) + first, next_state = causal_conv_sequence( + prompt, state, weight, indices, torch.tensor([0, 32, 64]) + ) + response = torch.cat((x[1056:], x[32:1024])) + second, next_state = causal_conv_sequence( + response, next_state, weight, indices.flip(0), torch.tensor([0, 32, 1024]) + ) + restored = torch.cat((first[:32], second[32:], first[32:], second[:32])) + exact(restored, whole) + exact(next_state, whole_state) + + +def test_response_gradient_reaches_prompt_and_convolution_weight(): + x, state, weight, indices, _ = inputs((8,)) + x.requires_grad_() + state.requires_grad_() + weight.requires_grad_() + _, prompt_state = causal_conv_sequence(x[:4], state, weight, indices, torch.tensor([0, 4])) + out, final_state = causal_conv_sequence( + x[4:], prompt_state, weight, indices, torch.tensor([0, 4]) + ) + (out.float().square().sum() + final_state.float().square().sum() * 0.01).backward() + for tensor in (x, weight): + assert tensor.grad is not None and torch.isfinite(tensor.grad).all() + assert tensor.grad.abs().max() > 0 + assert x.grad[:4].abs().max() > 0 + + +def test_empty_sequence_preserves_cache_and_gradients(): + x, state, weight, indices, cu = inputs((0,)) + state.requires_grad_() + out, next_state = causal_conv_sequence(x, state, weight, indices, cu) + assert out.shape == x.shape + exact(next_state, state) + next_state.float().sum().backward() + exact(state.grad, torch.ones_like(state)) diff --git a/tests/models/qwen3_next/check_qwen3_next_core_matrix.py b/tests/models/qwen3_next/check_qwen3_next_core_matrix.py new file mode 100644 index 000000000..823471772 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_core_matrix.py @@ -0,0 +1,101 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Mandatory TP4-shape GDN state tests; no full-model claim or optional skips.""" + +import pytest +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_provider import ( + GDNProviderConfig, + GDNState, + shared_gdn_core, +) +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def _inputs(lengths): + assert torch.cuda.is_available(), "Required shared-provider CUDA tests cannot skip" + generator = torch.Generator(device="cuda").manual_seed(20261002) + + def rand(*shape, dtype=torch.bfloat16): + return torch.randn(*shape, device="cuda", dtype=dtype, generator=generator) * 0.1 + + tokens = sum(lengths) + values = [rand(tokens, 2048), rand(tokens, 8), rand(tokens, 8), rand(tokens, 8, 128)] + weights = [rand(8, dtype=torch.float32), rand(8, dtype=torch.float32), rand(2048, 4), rand(128)] + state = GDNState( + rand(len(lengths) + 2, 2048, 3), + rand(len(lengths) + 2, 8, 128, 128, dtype=torch.float32), + ) + indices = torch.arange(1, len(lengths) + 1, device="cuda", dtype=torch.int32) + return values, weights, state, indices + + +def _run(values, weights, state, indices, lengths): + boundaries = torch.tensor([0, *torch.tensor(lengths).cumsum(0).tolist()], dtype=torch.int64) + return shared_gdn_core( + *values, *weights, state, indices, boundaries, config=GDNProviderConfig(), num_k_heads=4 + ) + + +@pytest.mark.parametrize("batch", [1, 4]) +@pytest.mark.parametrize("length", [8, 64, 256, 1024]) +@torch.no_grad() +def test_full_chunked_reordered_shared_core(batch, length): + lengths = [length] * batch + values, weights, initial, indices = _inputs(lengths) + whole, final = _run(values, weights, initial, indices, lengths) + for cut in (1, 63, 64, 65, 100, 528, 576): + cut = min(cut, length) + state = initial + pieces = [[] for _ in lengths] + for start, end in ((0, cut), (cut, length)): + # Reorder the response-bearing step, preserving sequence/cache identity. + order = list(range(batch)) if start == 0 else list(reversed(range(batch))) + rows = [row for seq in order for row in range(seq * length + start, seq * length + end)] + row_ids = torch.tensor(rows, device="cuda", dtype=torch.int64) + chunk = [value.index_select(0, row_ids) for value in values] + part, state = _run(chunk, weights, state, indices[order], [end - start] * batch) + for position, seq in enumerate(order): + pieces[seq].append(part[position * (end - start) : (position + 1) * (end - start)]) + exact(torch.cat([torch.cat(seq) for seq in pieces]), whole, name=f"output cut={cut}") + exact(state.convolution, final.convolution, name=f"conv state cut={cut}") + exact(state.recurrent, final.recurrent, name=f"recurrent state cut={cut}") + exact(final.convolution[0], initial.convolution[0], name="reserved conv slot") + exact(final.recurrent[0], initial.recurrent[0], name="reserved recurrent slot") + + +@torch.no_grad() +def test_variable_lengths_empty_and_mixed_step_match_independent_sequences(): + lengths = [65, 1, 0, 17] + values, weights, initial, indices = _inputs(lengths) + mixed, mixed_state = _run(values, weights, initial, indices, lengths) + start = 0 + for seq, length in enumerate(lengths): + end = start + length + output, state = _run( + [value[start:end] for value in values], + weights, + initial, + indices[seq : seq + 1], + [length], + ) + exact(output, mixed[start:end], name=f"sequence {seq}") + exact(state.convolution[seq + 1], mixed_state.convolution[seq + 1], name="conv state") + exact(state.recurrent[seq + 1], mixed_state.recurrent[seq + 1], name="recurrent state") + start = end + exact(mixed_state.recurrent[3], initial.recurrent[3], name="empty sequence state") + + +@pytest.mark.parametrize("batch", [1, 4]) +@torch.no_grad() +def test_shared_core_one_hundred_repetitions(batch): + lengths = [8] * batch + values, weights, initial, indices = _inputs(lengths) + expected, expected_state = _run(values, weights, initial, indices, lengths) + for repeat in range(100): + output, state = _run(values, weights, initial, indices, lengths) + exact(output, expected, name=f"output repeat={repeat}") + exact(state.convolution, expected_state.convolution, name="conv state") + exact(state.recurrent, expected_state.recurrent, name="recurrent state") diff --git a/tests/models/qwen3_next/check_qwen3_next_forward.py b/tests/models/qwen3_next/check_qwen3_next_forward.py new file mode 100644 index 000000000..29ae44382 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_forward.py @@ -0,0 +1,238 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Mandatory shared CUDA GEMM and MoE checks; these do not claim model L2.""" + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.models.qwen3_next.qwen3_next_forward import ( + combine_routes, + shared_linear, + shared_moe, + stable_top10_routes, +) +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def random(shape, seed=801, *, grad=False, dtype=torch.bfloat16): + assert torch.cuda.is_available(), "Mandatory shared-forward CUDA gates cannot skip" + gen = torch.Generator(device="cuda").manual_seed(seed) + return ( + torch.randn(shape, generator=gen, device="cuda", dtype=dtype) + .mul_(0.02) + .requires_grad_(grad) + ) + + +@pytest.mark.parametrize( + "k,n", [(2048, 3072), (2048, 16), (1024, 2048), (2048, 256), (128, 2048), (2048, 37984)] +) +@torch.no_grad() +def test_qwen_linear_shapes_batch_chunk_and_reorder(k, n): + x, weight = random((17, k)), random((n, k), seed=802) + output = shared_linear(x, weight) + exact( + torch.cat( + [ + shared_linear(x[:1], weight), + shared_linear(x[1:4], weight), + shared_linear(x[4:], weight), + ] + ), + output, + ) + order = torch.arange(16, -1, -1, device="cuda") + exact(shared_linear(x[order], weight), output[order]) + + +def test_shared_linear_vjp_and_empty_tokens(): + x, weight, bias = ( + random((4, 2048), grad=True), + random((16, 2048), 802, grad=True), + random((16,), 803, grad=True), + ) + output = shared_linear(x, weight, bias) + grad = random(output.shape, 804) + output.backward(grad) + refs = [value.detach().double().requires_grad_() for value in (x, weight, bias)] + F.linear(*refs).backward(grad.double()) + for value, ref in zip((x, weight, bias), refs): + torch.testing.assert_close(value.grad.float(), ref.grad.float(), atol=0.002, rtol=0.016) + assert (value.grad.double() - ref.grad).norm() / ref.grad.norm() < 0.016 + with torch.no_grad(): + exact(output, shared_linear(x, weight, bias)) + empty = random((0, 2048), grad=True) + weight.grad = bias.grad = None + shared_linear(empty, weight, bias).sum().backward() + assert empty.grad.shape == empty.shape + assert torch.count_nonzero(weight.grad) == 0 and torch.count_nonzero(bias.grad) == 0 + + +def test_stable_route_tie_boundary_and_router_gradient(): + logits = random((4, 512), dtype=torch.float32) + logits.zero_().requires_grad_() + routes = stable_top10_routes(logits) + assert torch.equal(routes.indices, torch.arange(10, device="cuda").expand(4, 10)) + exact(routes.weights, torch.full((4, 10), 0.1, device="cuda")) + (routes.weights * torch.arange(10, device="cuda")).sum().backward() + assert torch.isfinite(logits.grad).all() and logits.grad.abs().max() > 0 + bad = logits.detach().clone() + bad[0, 0] = torch.nan + with pytest.raises(ValueError, match="finite"): + stable_top10_routes(bad) + + +def test_combine_has_fixed_route_order_and_differentiable_weights(): + outputs = random((4, 10, 2048), grad=True) + weights = random((4, 10), 802, dtype=torch.float32, grad=True) + actual = combine_routes(outputs, weights) + exact( + torch.cat([combine_routes(outputs[i : i + 1], weights[i : i + 1]) for i in range(4)]), + actual, + ) + actual.float().square().sum().backward() + assert weights.grad.abs().max() > 0 and outputs.grad.abs().max() > 0 + refs = [value.detach().double().requires_grad_() for value in (outputs, weights)] + expected = (refs[0] * refs[1].unsqueeze(-1)).sum(dim=1) + expected.backward((2 * actual.float()).double()) + for value, reference in zip((outputs, weights), refs): + torch.testing.assert_close( + value.grad.float(), reference.grad.float(), atol=0.002, rtol=0.016 + ) + + +def test_qwen_moe_actual_tp4_shapes_and_independent_training_forward(): + x = random((4, 2048), grad=True) + router = random((512, 2048), 802, grad=True) + gate_up = random((512, 256, 2048), 803, grad=True) + down = random((512, 2048, 128), 804, grad=True) + output, routes = shared_moe(x, router, gate_up, down) + with torch.no_grad(): + row_output, row_routes = shared_moe(x[:1], router, gate_up, down) + exact(row_output, output[:1]) + exact(row_routes.weights, routes.weights[:1]) + assert torch.equal(row_routes.indices, routes.indices[:1]) + output.float().square().sum().backward() + for value in (x, router, gate_up, down): + assert value.grad is not None and torch.isfinite(value.grad).all() + assert value.grad.abs().max() > 0 + active = routes.indices.unique() + inactive = torch.ones(512, dtype=torch.bool, device="cuda") + inactive[active] = False + assert torch.count_nonzero(gate_up.grad[inactive]) == 0 + assert torch.count_nonzero(down.grad[inactive]) == 0 + + +def test_moe_vjp_matches_independent_double_formula(): + # A small hidden width isolates the derivative; the previous gate covers + # the unmodified official TP4 widths and expert count. + x = random((2, 32)).mul_(10).requires_grad_() + router = random((512, 32), 802).mul_(5).requires_grad_() + gate_up = random((512, 16, 32), 803).mul_(10).requires_grad_() + down = random((512, 32, 8), 804).mul_(10).requires_grad_() + actual, routes = shared_moe(x, router, gate_up, down) + grad = random(actual.shape, 805) + actual.backward(grad) + xr, rr, gu, dr = [v.detach().double().requires_grad_() for v in (x, router, gate_up, down)] + probability = (xr @ rr.t()).softmax(-1) + selected = probability.gather(1, routes.indices) + selected = selected / selected.sum(-1, keepdim=True) + rows = [] + for token in range(2): + slots = [] + for slot, expert in enumerate(routes.indices[token].tolist()): + gate, up = F.linear(xr[token], gu[expert]).chunk(2) + slots.append(F.linear(F.silu(gate) * up, dr[expert]) * selected[token, slot]) + rows.append(torch.stack(slots).sum(0)) + torch.stack(rows).backward(grad.double()) + for value, reference in zip((x, router, gate_up, down), (xr, rr, gu, dr)): + assert torch.isfinite(value.grad).all() + # BF16 intermediates are rounded by the actual provider; this compares + # derivative correctness, never the raw-bit L1/L2 acceptance criterion. + torch.testing.assert_close( + value.grad.float(), reference.grad.float(), atol=0.00002, rtol=0.016 + ) + assert (value.grad.double() - reference.grad).norm() / reference.grad.norm() < 0.016 + + +@torch.no_grad() +def test_route_weights_are_batch_invariant(): + logits = random((32, 512), 805, dtype=torch.float32).mul_(100) + whole = stable_top10_routes(logits) + for count in (1, 8, 17): + part = stable_top10_routes(logits[:count]) + assert torch.equal(part.indices, whole.indices[:count]) + exact(part.weights, whole.weights[:count]) + moe_x, router = random((32, 2048), 806), random((512, 2048), 807) + gate_up, down = random((512, 256, 2048), 808), random((512, 2048, 128), 809) + whole_out, _ = shared_moe(moe_x, router, gate_up, down) + part_out, _ = shared_moe(moe_x[:8], router, gate_up, down) + exact(part_out, whole_out[:8]) + + +class _PerExpertReference(torch.autograd.Function): + """The previous provider: one pinned vLLM GEMM pair per expert, in expert order.""" + + @staticmethod + def forward(ctx, x, gate_up, down, route_indices): + from rl_engine.models.qwen3_next.qwen3_next_forward import _vllm_linear + + output = x.new_empty((x.shape[0] * 10, x.shape[1])) + flat = route_indices.flatten() + for expert in flat.unique().tolist(): + slots = (flat == expert).nonzero().squeeze(1) + rows = x.index_select(0, slots // 10) + gate, up = _vllm_linear()(rows, gate_up[expert]).chunk(2, dim=-1) + activated = (F.silu(gate.float()) * up.float()).to(x.dtype) + output.index_copy_(0, slots, _vllm_linear()(activated, down[expert])) + return output.reshape(x.shape[0], 10, x.shape[1]) + + +@torch.no_grad() +@pytest.mark.parametrize("width", [128, 512]) +@pytest.mark.parametrize("tokens", [1, 7, 13, 300]) +def test_grouped_experts_are_bitwise_the_per_expert_gemms(width, tokens): + from rl_engine.models.qwen3_next.qwen3_next_forward import _RoutedExperts, shared_router + + x, router = random((tokens, 2048), 811), random((512, 2048), 812) + gate_up, down = random((512, 2 * width, 2048), 813), random((512, 2048, width), 814) + indices = stable_top10_routes(shared_router(x, router)).indices + exact( + _RoutedExperts.apply(x, gate_up, down, indices), + _PerExpertReference.apply(x, gate_up, down, indices), + ) + + +@torch.no_grad() +@pytest.mark.parametrize("width", [128, 512]) +def test_grouped_projection_is_bitwise_the_per_expert_gemm(width): + # The backward recomputes [gate, up] with the grouped kernel; equality here + # leaves every gradient bitwise what the per-expert backward produced. + from rl_engine.models.qwen3_next.qwen3_next_forward import ( + _vllm_linear, + grouped_route_linear, + shared_router, + ) + + x, router = random((300, 2048), 817), random((512, 2048), 818) + gate_up = random((512, 2 * width, 2048), 819) + indices = stable_top10_routes(shared_router(x, router)).indices + grouped = grouped_route_linear(x, gate_up, indices).reshape(-1, 2 * width) + flat = indices.flatten() + for expert in flat.unique().tolist(): + slots = (flat == expert).nonzero().squeeze(1) + exact(grouped[slots], _vllm_linear()(x[slots // 10], gate_up[expert]), name=f"e{expert}") + + +def test_grouped_experts_fail_closed_on_the_tensor_descriptor_path(monkeypatch): + from rl_engine.models.qwen3_next import qwen3_next_forward as provider + + prepare, dispatch, _ = provider._vllm_grouped() + monkeypatch.setattr(provider, "_vllm_grouped", lambda: (prepare, dispatch, lambda: True)) + x = random((4, 2048), 815) + with pytest.raises(RuntimeError, match="VLLM_TRITON_USE_TD"): + provider.grouped_route_linear( + x, random((512, 256, 2048), 816), torch.zeros(4, 10, device="cuda", dtype=torch.long) + ) diff --git a/tests/models/qwen3_next/check_qwen3_next_gdn_bridge.py b/tests/models/qwen3_next/check_qwen3_next_gdn_bridge.py new file mode 100644 index 000000000..9656b80c7 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_gdn_bridge.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Explicit subprocess provider tests; never import vLLM into ordinary pytest.""" + +import pytest +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_gdn import ( + _provider, + packed_decode_training_step, +) +from rl_engine.reference.linear_attn import GatedDeltaRuleRecurrentStepOp +from rl_engine.validation.common.tensor_identity import ( + assert_tensor_bitwise_equal, + tensor_bitwise_equal, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +def _inputs(seed=1): + gen = torch.Generator(device="cuda").manual_seed(seed) + + def rand(*shape, dtype=torch.bfloat16): + return torch.randn(*shape, device="cuda", dtype=dtype, generator=gen) + + return [ + rand(2, 512), + rand(2, 2), + rand(2, 2), + rand(2, dtype=torch.float32), + rand(2, dtype=torch.float32), + rand(3, 2, 128, 128, dtype=torch.float32) * 0.01, + ] + + +def test_bridge_forward_is_provider_exact_and_does_not_mutate_inputs(): + inputs = _inputs() + state_before = inputs[-1].clone() + indices = torch.tensor([1, 2], device="cuda", dtype=torch.int32) + out, state = packed_decode_training_step(*inputs, indices, num_k_heads=1) + expected = torch.empty_like(out) + expected_state = state_before.clone() + _provider()( + *inputs[:-1], 128**-0.5, expected_state, expected, indices, use_qk_l2norm_in_kernel=True + ) + assert_tensor_bitwise_equal(out, expected) + assert_tensor_bitwise_equal(state, expected_state) + assert tensor_bitwise_equal(inputs[-1], state_before) + + +def test_bridge_keeps_prompt_state_gradient_through_response(): + inputs = [x.requires_grad_() for x in _inputs()] + indices = torch.tensor([1, 2], device="cuda", dtype=torch.int32) + _, prompt_state = packed_decode_training_step(*inputs, indices, num_k_heads=1) + prompt_state.retain_grad() + response = _inputs(seed=2)[0].requires_grad_() + out, final_state = packed_decode_training_step( + response, *inputs[1:5], prompt_state, indices, num_k_heads=1 + ) + (out.float().square().sum() + final_state.square().sum() * 0.01).backward() + assert prompt_state.grad is not None and prompt_state.grad.abs().max() > 0 + for tensor in [*inputs, response]: + assert tensor.grad is not None + assert torch.isfinite(tensor.grad).all() + assert tensor.grad.abs().max() > 0 + + +def test_bridge_backward_matches_recomputed_recurrence_vjp(): + inputs = [x.requires_grad_() for x in _inputs()] + indices = torch.tensor([1, 2], device="cuda", dtype=torch.int32) + out, state = packed_decode_training_step(*inputs, indices, num_k_heads=1) + grad_out = torch.ones_like(out) * 0.125 + grad_state = torch.ones_like(state) * 0.001 + torch.autograd.backward((out, state), (grad_out, grad_state)) + refs = [x.detach().clone().requires_grad_() for x in inputs] + ref_out, ref_state = GatedDeltaRuleRecurrentStepOp()( + *refs, indices, scale=128**-0.5, num_k_heads=1 + ) + torch.autograd.backward((ref_out, ref_state), (grad_out, grad_state)) + for actual, ref in zip(inputs, refs): + assert tensor_bitwise_equal(actual.grad, ref.grad) + + +def test_initial_state_vjp_matches_provider_directional_difference(): + inputs = _inputs() + inputs[-1].requires_grad_() + indices = torch.tensor([1, 2], device="cuda", dtype=torch.int32) + _, output = packed_decode_training_step(*inputs, indices, num_k_heads=1) + direction = torch.ones_like(inputs[-1]) * 0.03125 + probe = torch.linspace(-0.01, 0.01, output.numel(), device="cuda").reshape_as(output) + (output.double() * probe.double()).sum().backward() + analytic = (inputs[-1].grad.double() * direction.double()).sum() + eps = 0.01 + with torch.no_grad(): + _, plus = packed_decode_training_step( + *inputs[:-1], inputs[-1] + eps * direction, indices, num_k_heads=1 + ) + _, minus = packed_decode_training_step( + *inputs[:-1], inputs[-1] - eps * direction, indices, num_k_heads=1 + ) + numerical = ((plus.double() - minus.double()) * probe.double()).sum() / (2 * eps) + torch.testing.assert_close(analytic, numerical, rtol=1e-3, atol=1e-4) + + +@pytest.mark.parametrize("empty", [False, True]) +def test_inactive_state_is_identity_in_backward(empty): + inputs = _inputs() + if empty: + inputs[:3] = [x[:0] for x in inputs[:3]] + inputs = [x.requires_grad_() for x in inputs] + indices = torch.zeros(inputs[0].shape[0], device="cuda", dtype=torch.int32) + out, state = packed_decode_training_step(*inputs, indices, num_k_heads=1) + (out.float().sum() + state.sum()).backward() + assert tensor_bitwise_equal(inputs[-1].grad, torch.ones_like(state)) + for tensor in inputs[:-1]: + assert torch.count_nonzero(tensor.grad) == 0 diff --git a/tests/models/qwen3_next/check_qwen3_next_gdn_sequence.py b/tests/models/qwen3_next/check_qwen3_next_gdn_sequence.py new file mode 100644 index 000000000..6fa53afb2 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_gdn_sequence.py @@ -0,0 +1,102 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Packed decode recurrence boundary checks; no complete-model L2 claim.""" + +import pytest +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_gdn import packed_recurrent_sequence +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def inputs(lengths, seed=7): + assert torch.cuda.is_available(), "CUDA is required; this acceptance suite cannot skip" + generator = torch.Generator(device="cuda").manual_seed(seed) + + def rand(*shape, dtype=torch.bfloat16): + return torch.randn(*shape, device="cuda", dtype=dtype, generator=generator) + + tokens = sum(lengths) + # Actual TP4 shard: 16 / 4 key heads, 32 / 4 value heads, K=V=128. + tensors = [ + rand(tokens, 2048), + rand(tokens, 8), + rand(tokens, 8), + rand(8, dtype=torch.float32), + rand(8, dtype=torch.float32), + rand(len(lengths) + 1, 8, 128, 128, dtype=torch.float32) * 0.01, + ] + indices = torch.arange(1, len(lengths) + 1, device="cuda", dtype=torch.int32) + cu = torch.tensor([0, *torch.tensor(lengths).cumsum(0).tolist()], dtype=torch.int64) + return tensors, indices, cu + + +def run(tensors, indices, cu): + return packed_recurrent_sequence(*tensors, indices, cu, num_k_heads=4) + + +@torch.no_grad() +@pytest.mark.parametrize("lengths", [(8,), (64, 8, 0, 16), (1024, 256, 64, 8)]) +def test_packed_matches_independent_sequences_and_reordering(lengths): + tensors, indices, cu = inputs(lengths) + original_state = tensors[-1].clone() + output, state = run(tensors, indices, cu) + for seq, length in enumerate(lengths): + start, end = cu[seq : seq + 2].tolist() + one = [*(value[start:end] for value in tensors[:3]), *tensors[3:]] + expected, expected_state = run(one, indices[seq : seq + 1], torch.tensor([0, length])) + exact(output[start:end], expected) + exact(state[seq + 1], expected_state[seq + 1]) + permutation = list(reversed(range(len(lengths)))) + rows = torch.cat([torch.arange(cu[seq], cu[seq + 1], device="cuda") for seq in permutation]) + changed = [*(value.index_select(0, rows) for value in tensors[:3]), *tensors[3:]] + reordered_cu = torch.tensor([0, *torch.tensor([lengths[i] for i in permutation]).cumsum(0)]) + reordered, reordered_state = run(changed, indices[permutation], reordered_cu) + exact(reordered, output.index_select(0, rows)) + exact(reordered_state, state) + exact(tensors[-1], original_state) + + +@torch.no_grad() +def test_1024_token_chunk_continuation_is_bitwise_exact(): + tensors, indices, cu = inputs((1024,)) + whole, whole_state = run(tensors, indices, cu) + state, pieces = tensors[-1], [] + cuts = (0, 1, 8, 64, 256, 513, 1024) + for start, end in zip(cuts, cuts[1:]): + chunk = [*(value[start:end] for value in tensors[:3]), *tensors[3:5], state] + output, state = run(chunk, indices, torch.tensor([0, end - start])) + pieces.append(output) + exact(torch.cat(pieces), whole) + exact(state, whole_state) + + +def test_chunk_boundary_retains_prompt_and_initial_state_gradients(): + tensors, indices, cu = inputs((8,)) + tensors = [tensor.requires_grad_() for tensor in tensors] + prompt = [*(value[:4] for value in tensors[:3]), *tensors[3:]] + _, prompt_state = run(prompt, indices, torch.tensor([0, 4])) + prompt_state.retain_grad() + response = [*(value[4:] for value in tensors[:3]), *tensors[3:5], prompt_state] + out, state = run(response, indices, torch.tensor([0, 4])) + (out.float().square().sum() + state.square().sum() * 0.01).backward() + assert prompt_state.grad is not None and prompt_state.grad.abs().max() > 0 + for value in tensors: + assert value.grad is not None and torch.isfinite(value.grad).all() + assert value.grad.abs().max() > 0 + assert tensors[0].grad[:4].abs().max() > 0 + + +@pytest.mark.parametrize("bad", [[1, 8], [0, 9], [0, 6, 5, 8]]) +def test_bad_packing_is_rejected(bad): + tensors, indices, _ = inputs((8,)) + with pytest.raises(ValueError, match="cu_seqlens"): + run(tensors, indices, torch.tensor(bad)) + + +def test_reserved_and_duplicate_slots_are_rejected(): + tensors, indices, cu = inputs((4, 4)) + for bad in ([0, 1], [1, 1]): + with pytest.raises(ValueError): + run(tensors, torch.tensor(bad, device="cuda", dtype=torch.int32), cu) 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/check_qwen3_next_shared_core.py b/tests/models/qwen3_next/check_qwen3_next_shared_core.py new file mode 100644 index 000000000..f850b0bee --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_shared_core.py @@ -0,0 +1,87 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Combined provider contract; does not instantiate a complete Qwen model.""" + +from dataclasses import replace + +import pytest +import torch + +from rl_engine.integrations.engines.train.vllm.qwen3_next_provider import ( + GDNProviderConfig, + GDNState, + shared_gdn_core, +) +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def inputs(tokens, grad=False): + assert torch.cuda.is_available(), "Required CUDA provider suite cannot skip" + gen = torch.Generator(device="cuda").manual_seed(73) + + def rand(*shape, dtype=torch.bfloat16): + return torch.randn(*shape, device="cuda", dtype=dtype, generator=gen).requires_grad_(grad) + + args = [ + rand(tokens, 2048), + rand(tokens, 8), + rand(tokens, 8), + rand(tokens, 8, 128), + rand(8, dtype=torch.float32), + rand(8, dtype=torch.float32), + rand(2048, 4), + rand(128), + GDNState(rand(2, 2048, 3), rand(2, 8, 128, 128, dtype=torch.float32)), + ] + return args, torch.tensor([1], device="cuda", dtype=torch.int32) + + +@pytest.mark.parametrize("tokens", [8, 1024]) +@torch.no_grad() +def test_shared_core_chunk_and_decode_state_match(tokens): + args, indices = inputs(tokens) + config = GDNProviderConfig() + output, state = shared_gdn_core( + *args, indices, torch.tensor([0, tokens]), config=config, num_k_heads=4 + ) + pieces, next_state = [], args[-1] + for start, end in ((0, 3), (3, tokens - 1), (tokens - 1, tokens)): + chunk = [*(value[start:end] for value in args[:4]), *args[4:8], next_state] + part, next_state = shared_gdn_core( + *chunk, indices, torch.tensor([0, end - start]), config=config, num_k_heads=4 + ) + pieces.append(part) + exact(torch.cat(pieces), output) + exact(next_state.convolution, state.convolution) + exact(next_state.recurrent, state.recurrent) + + +def test_independent_training_forward_matches_inference_and_retains_gradients(): + args, indices = inputs(8, grad=True) + config = GDNProviderConfig() + with torch.no_grad(): + reference, _ = shared_gdn_core( + *args, indices, torch.tensor([0, 8]), config=config, num_k_heads=4 + ) + prompt = [*(value[:4] for value in args[:4]), *args[4:]] + _, prompt_state = shared_gdn_core( + *prompt, indices, torch.tensor([0, 4]), config=config, num_k_heads=4 + ) + prompt_state.recurrent.retain_grad() + response = [*(value[4:] for value in args[:4]), *args[4:8], prompt_state] + actual, _ = shared_gdn_core( + *response, indices, torch.tensor([0, 4]), config=config, num_k_heads=4 + ) + exact(actual, reference[4:]) + actual.float().square().sum().backward() + for value in args[:8]: + assert value.grad is not None and torch.isfinite(value.grad).all() + assert value.grad.abs().max() > 0 + assert args[0].grad[:4].abs().max() > 0 + assert prompt_state.recurrent.grad.abs().max() > 0 + + +def test_unimplemented_provider_profile_is_rejected(): + with pytest.raises(ValueError, match="Unsupported"): + replace(GDNProviderConfig(), recurrence="unverified-prefill").validate() 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_forward_contract.py b/tests/models/qwen3_next/test_qwen3_next_forward_contract.py new file mode 100644 index 000000000..8259f9784 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_forward_contract.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CPU checks of the strict shared-forward boundary, without a CPU fallback.""" + +import pytest +import torch + +from rl_engine.models.qwen3_next import qwen3_next_forward as provider + + +@pytest.mark.parametrize( + "call", + [ + lambda: provider.shared_linear(torch.zeros(1, 8), torch.zeros(8, 8)), + lambda: provider.stable_top10_routes(torch.zeros(1, 512)), + lambda: provider.combine_routes(torch.zeros(1, 10, 8), torch.zeros(1, 10)), + lambda: provider.shared_attention(*(torch.zeros(1, 4, 8, 256) for _ in range(3))), + ], +) +def test_required_cuda_primitives_reject_cpu(call): + with pytest.raises(ValueError, match="no CPU fallback"): + call() + + +def test_shared_linear_rejects_an_unpinned_vllm(monkeypatch): + provider._vllm_linear.cache_clear() + monkeypatch.setattr(provider, "version", lambda package: "0.30.1") + with pytest.raises(RuntimeError, match="pinned vLLM 0.30.0"): + provider._vllm_linear() + provider._vllm_linear.cache_clear() + + +def _rows(count, width, seed): + generator = torch.Generator().manual_seed(seed) + return torch.randn((count, width), generator=generator, dtype=torch.float32) * 4 + + +@pytest.mark.parametrize("width", [1, 2, 7, 10, 512]) +def test_fixed_order_row_sum_matches_sequential_double_sum(width): + values = _rows(5, width, 901) + actual = provider.fixed_order_row_sum(values) + assert actual.shape == (5, 1) + reference = values.double().sum(-1, keepdim=True).float() + torch.testing.assert_close(actual, reference, rtol=1e-5, atol=1e-4) + + +@pytest.mark.parametrize("width", [10, 512]) +def test_fixed_order_row_sum_is_independent_of_row_count(width): + values = _rows(32, width, 902) + whole = provider.fixed_order_row_sum(values) + for count in (1, 8, 17): + assert torch.equal(provider.fixed_order_row_sum(values[:count]), whole[:count]) + + +def test_fixed_order_softmax_rows_do_not_depend_on_batch(): + logits = _rows(32, 512, 903) + whole = provider.fixed_order_softmax(logits) + torch.testing.assert_close(whole.double().sum(-1), torch.ones(32, dtype=torch.float64)) + for count in (1, 8): + assert torch.equal(provider.fixed_order_softmax(logits[:count]), whole[:count]) + # The row sum is a fixed sequential chain, so each result is a pure function of + # its own row and cannot depend on how many rows share the call. 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_tp_blocks.py b/tests/models/qwen3_next/test_qwen3_next_tp_blocks.py new file mode 100644 index 000000000..14de47b97 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_tp_blocks.py @@ -0,0 +1,162 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""HF mappings, KV replica ownership and official Qwen3-Next attention/MoE dimensions.""" + +import pytest +import torch + +from rl_engine.models.qwen3_next import qwen3_next_tp +from rl_engine.models.qwen3_next import qwen3_next_tp_blocks as blocks +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal as exact + + +def pattern(shape): + return ( + torch.arange(torch.Size(shape).numel(), dtype=torch.int32) + .remainder(251) + .to(torch.bfloat16) + .reshape(shape) + ) + + +@pytest.fixture(scope="module") +def attention_weights(): + return {name: pattern(shape) for name, shape in blocks.ATTENTION_HF_SHAPES.items()} + + +def test_attention_hf_roundtrip_and_kv_pair_layout(attention_weights): + shards = [blocks.shard_attention_weights(attention_weights, rank) for rank in range(4)] + for name, value in blocks.assemble_attention_weights(shards).items(): + exact(value, attention_weights[name], name=name) + for rank, shard in enumerate(shards): + exact( + shard["q_proj.weight"], + attention_weights["q_proj.weight"][rank * 2048 : (rank + 1) * 2048], + ) + begin = rank // 2 * 256 + exact(shard["k_proj.weight"], attention_weights["k_proj.weight"][begin : begin + 256]) + + +def test_attention_export_rejects_disagreeing_kv_replica(attention_weights): + shards = [blocks.shard_attention_weights(attention_weights, rank) for rank in range(4)] + shards[1]["k_proj.weight"] = shards[1]["k_proj.weight"].clone() + shards[1]["k_proj.weight"][0, 0] = -1 + with pytest.raises(ValueError, match="Replicated KV"): + blocks.assemble_attention_weights(shards) + + +@pytest.mark.parametrize( + "name", + [ + "gate.weight", + "shared_expert_gate.weight", + "shared_expert.gate_proj.weight", + "shared_expert.up_proj.weight", + "shared_expert.down_proj.weight", + "experts.0.gate_proj.weight", + "experts.0.up_proj.weight", + "experts.511.down_proj.weight", + ], +) +def test_moe_hf_weight_roundtrip(name): + value = pattern(blocks.moe_hf_shapes()[name]) + shards = [blocks.shard_moe_weight(name, value, rank) for rank in range(4)] + exact(blocks.assemble_moe_weight(name, shards), value, name=name) + if name.endswith("down_proj.weight"): + assert shards[0].shape == (2048, 128) + elif name.endswith(("gate_proj.weight", "up_proj.weight")): + assert shards[0].shape == (128, 2048) + + +def test_moe_export_rejects_router_replica_drift(): + values = [pattern((512, 2048)) for _ in range(4)] + values[2][0, 0] = -1 + with pytest.raises(ValueError, match="Replicated MoE"): + blocks.assemble_moe_weight("gate.weight", values) + + +@pytest.mark.parametrize("rank", [-1, 4, True, 0.5]) +def test_invalid_shard_rank_rejected(rank): + with pytest.raises(ValueError, match="TP4 rank"): + blocks.shard_moe_weight("gate.weight", torch.empty(512, 2048, dtype=torch.bfloat16), rank) + + +def test_official_shared_expert_width_is_512_and_each_expert_has_three_weights(): + shapes = blocks.moe_hf_shapes() + assert len(shapes) == 512 * 3 + 5 + assert shapes["shared_expert.gate_proj.weight"] == (512, 2048) + assert shapes["shared_expert.down_proj.weight"] == (2048, 512) + + +def test_tp4_moe_rejects_a_missing_or_wrong_size_tp_group(monkeypatch): + monkeypatch.setattr(qwen3_next_tp.dist, "is_initialized", lambda: False) + with pytest.raises(ValueError, match="four-rank TP group"): + blocks.TP4MoE(group=None, device="meta") + monkeypatch.setattr(qwen3_next_tp.dist, "is_initialized", lambda: True) + monkeypatch.setattr(qwen3_next_tp.dist, "get_world_size", lambda group: 8) + with pytest.raises(ValueError, match="four-rank TP group"): + blocks.TP4MoE(group=None, device="meta") + + +def test_tp4_moe_parameter_ownership(monkeypatch): + group = object() + monkeypatch.setattr(qwen3_next_tp.dist, "is_initialized", lambda: True) + monkeypatch.setattr(qwen3_next_tp.dist, "get_world_size", lambda group: 4) + monkeypatch.setattr(blocks.dist, "get_rank", lambda group: 2) + module = blocks.TP4MoE(group=group, device="meta") + assert module.rank == 2 + assert module.experts.gate_up.shape == (512, 256, 2048) + assert module.experts.down.shape == (512, 2048, 128) + assert ( + module.experts.gate_up.partition_dim == 1 and module.experts.gate_up.partition_stride == 2 + ) + assert module.experts.down.partition_dim == 2 + assert module.shared_expert.down_proj.weight.partition_dim == 1 + # Router and shared-expert gate are replicated: no TP attributes, full shapes. + assert not getattr(module.gate.weight, "tensor_model_parallel", False) + assert not getattr(module.shared_expert_gate.weight, "tensor_model_parallel", False) + assert module.gate.weight.shape == (512, 2048) + + +@pytest.fixture +def fake_tp(monkeypatch): + group = object() + monkeypatch.setattr(blocks.dist, "is_initialized", lambda: True) + monkeypatch.setattr(blocks.dist, "get_world_size", lambda group: 4) + monkeypatch.setattr(blocks, "_kv_replica_groups", lambda group: (object(), object())) + return group + + +@pytest.mark.parametrize("rank", [0, 1, 2, 3]) +def test_attention_parameter_ownership_and_actual_hf_storage( + fake_tp, monkeypatch, rank, attention_weights +): + monkeypatch.setattr(blocks.dist, "get_rank", lambda group: rank) + module = blocks.TP4FullAttention(group=fake_tp, device="cpu") + module.load_hf_weights(attention_weights) + expected = blocks.shard_attention_weights(attention_weights, rank) + for name, value in module.export_local_hf_weights().items(): + exact(value, expected[name], name=name) + assert module.q_proj.weight.tensor_model_parallel + assert module.o_proj.weight.partition_dim == 1 + assert module.k_proj.weight.shared == (rank % 2 == 1) + assert module.v_proj.weight.shared == (rank % 2 == 1) + assert not getattr(module.q_norm.weight, "tensor_model_parallel", False) + state = module.initial_state(3) + assert len(state.keys) == len(state.values) == 3 + assert all(value.shape == (1, 1, 0, 256) for value in state.keys) + + +def test_rope_uses_absolute_positions_and_preserves_unrotated_tail(fake_tp, monkeypatch): + monkeypatch.setattr(blocks.dist, "get_rank", lambda group: 0) + module = blocks.TP4FullAttention(group=fake_tp, device="cpu") + x = torch.randn(8, 4, 256, dtype=torch.bfloat16) + positions = torch.arange(61, 69) + whole = module._rotate(x, positions) + exact(whole[..., 64:], x[..., 64:]) + exact( + torch.cat([module._rotate(x[:3], positions[:3]), module._rotate(x[3:], positions[3:])]), + whole, + ) + assert not torch.equal(module._rotate(x, torch.arange(8)), whole) diff --git a/tests/models/qwen3_next/test_qwen3_next_tp_gdn.py b/tests/models/qwen3_next/test_qwen3_next_tp_gdn.py new file mode 100644 index 000000000..953472169 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_tp_gdn.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Official HF interleaving and TP shard round trips, without a GPU dependency.""" + +import pytest +import torch + +from rl_engine.models.qwen3_next.qwen3_next_tp_gdn import ( + _GLOBAL_SHAPES, + assemble_gdn_weights, + shard_gdn_weights, +) +from rl_engine.validation.common.tensor_identity import assert_tensor_bitwise_equal + + +@pytest.fixture(scope="module") +def weights(): + return { + name: torch.arange(torch.Size(shape).numel(), dtype=torch.int32) + .remainder(251) + .to(torch.bfloat16) + .reshape(shape) + for name, shape in _GLOBAL_SHAPES.items() + } + + +def test_official_interleaving_survives_tp4_round_trip(weights): + shards = [shard_gdn_weights(weights, rank) for rank in range(4)] + restored = assemble_gdn_weights(shards) + for name, tensor in weights.items(): + assert_tensor_bitwise_equal(restored[name], tensor, name=name) + for rank, shard in enumerate(shards): + # Conv is Q/K/V segmented, unlike the GQA-interleaved qkvz projection. + assert_tensor_bitwise_equal( + shard["conv1d.weight"][:512], weights["conv1d.weight"][rank * 512 : (rank + 1) * 512] + ) + assert_tensor_bitwise_equal( + shard["conv1d.weight"][512:1024], + weights["conv1d.weight"][2048 + rank * 512 : 2048 + (rank + 1) * 512], + ) + assert_tensor_bitwise_equal( + shard["in_proj_qkvz.weight"], + weights["in_proj_qkvz.weight"][rank * 3072 : (rank + 1) * 3072], + ) + + +@pytest.mark.parametrize("rank", [-1, 4, True, 1.0]) +def test_invalid_tp_rank_is_rejected(weights, rank): + with pytest.raises(ValueError, match="TP4 rank"): + shard_gdn_weights(weights, rank) + + +def test_divergent_replicated_norm_is_not_silently_exported(weights): + shards = [shard_gdn_weights(weights, rank) for rank in range(4)] + shards[1]["norm.weight"] = shards[1]["norm.weight"].clone() + shards[1]["norm.weight"][0] = -1 + with pytest.raises(ValueError, match="Replicated"): + assemble_gdn_weights(shards) + + +def test_missing_hf_weight_is_rejected(weights): + incomplete = {name: value for name, value in weights.items() if name != "dt_bias"} + with pytest.raises(ValueError, match="weight names"): + shard_gdn_weights(incomplete, 0) 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/common/test_tensor_identity.py b/tests/validation/common/test_tensor_identity.py new file mode 100644 index 000000000..dd3a1ed13 --- /dev/null +++ b/tests/validation/common/test_tensor_identity.py @@ -0,0 +1,44 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CPU regressions for the raw-bit identity used by Qwen3-Next acceptance checks.""" + +import pytest +import torch + +from rl_engine.validation.common.tensor_identity import ( + assert_tensor_bitwise_equal, + tensor_bitwise_equal, +) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +def test_signed_zero_is_not_identical(dtype): + positive = torch.tensor([0.0], dtype=dtype) + negative = torch.tensor([-0.0], dtype=dtype) + assert not tensor_bitwise_equal(positive, negative) + with pytest.raises(AssertionError, match="raw-bit"): + assert_tensor_bitwise_equal(positive, negative) + + +@pytest.mark.parametrize("value", [float("nan"), float("inf"), -float("inf")]) +def test_identical_nonfinite_values_are_rejected(value): + tensor = torch.tensor([value]) + assert not tensor_bitwise_equal(tensor, tensor) + + +def test_same_numbers_in_different_dtypes_are_rejected(): + assert not tensor_bitwise_equal(torch.ones(2), torch.ones(2, dtype=torch.bfloat16)) + + +def test_adjacent_bfloat16_values_are_rejected(): + value = torch.ones(2, dtype=torch.bfloat16) + changed = (value.view(torch.int16) + 1).view(torch.bfloat16) + assert not tensor_bitwise_equal(value, changed) + + +@pytest.mark.parametrize( + "value", [torch.tensor(1.0), torch.empty(0), torch.arange(12).reshape(3, 4).T] +) +def test_logical_identity_does_not_require_matching_strides(value): + assert_tensor_bitwise_equal(value, value.contiguous().clone()) 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_attention_prior_art.py b/tools/validation/models/plot_qwen3_next_attention_prior_art.py new file mode 100644 index 000000000..6a3fd7a3f --- /dev/null +++ b/tools/validation/models/plot_qwen3_next_attention_prior_art.py @@ -0,0 +1,152 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Render the TP4 attention prior-art report as a PNG. + + python tools/validation/models/plot_qwen3_next_attention_prior_art.py \\ + docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/report.json + +writes ``figure.png`` next to the report: which invariance checks each +implementation passes, accuracy against FP64, and prefill/decode latency. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 +from matplotlib.colors import ListedColormap # noqa: E402 + +RL_KERNEL = "#2a78d6" +INK = "#2b2b29" +MUTED = "#6f6e69" +GRID = "#e4e3dd" +SURFACE = "#fcfcfb" +PASS = "#1baf7a" +FAIL = "#d64545" +NA = "#e4e3dd" +EXISTING = ( + "#eb6834", + "#b8501f", + "#f29a62", + "#8a3a15", + "#f6bd94", + "#d9773f", + "#6b2a0e", + "#c96a2c", + "#f0a878", +) + +LABELS = { + "rl_kernel_cuda": "RL-Kernel deterministic", + "torch_sdpa": "torch SDPA", + "vllm_fa2_auto": "vLLM FA2 (num_splits auto)", + "vllm_fa2_split1": "vLLM FA2 (num_splits=1)", + "vllm_triton_2d": "vLLM Triton unified (2D)", + "flashinfer": "FlashInfer ragged prefill", + "fa4_cute": "FlashAttention-4 (CuTe)", + "te_fused": "TE, cuDNN fused (inference)", + "te_training": "TE, unfused (training)", + "megatron_local": "Megatron-core local", +} +COLUMNS = ( + ("batch_bitwise", "first", "batch: first"), + ("batch_bitwise", "last", "batch: last"), + ("prefill_decode_bitwise", "last_64", "chunk: last 64"), + ("prefill_decode_bitwise", "last_1", "decode: last 1"), + ("backward_batch_bitwise", "dq", "bwd dq (packed)"), + ("backward_batch_bitwise", "dk", "bwd dk (packed)"), + ("backward_batch_bitwise", "dv", "bwd dv (packed)"), +) + + +def _style(ax): + ax.set_facecolor(SURFACE) + for side in ("top", "right"): + ax.spines[side].set_visible(False) + ax.tick_params(colors=INK, labelsize=8) + ax.grid(color=GRID, linewidth=0.6) + + +def plot(report: dict, out: Path) -> None: + cands = [c for c in report["candidates"] if "unavailable" not in c and "failed" not in c] + keys = [c["key"] for c in cands] + existing = iter(EXISTING) + colors = {k: RL_KERNEL if k == "rl_kernel_cuda" else next(existing) for k in keys} + labels = [LABELS.get(k, k) for k in keys] + fig, axes = plt.subplots(1, 3, figsize=(17, 4.6), gridspec_kw={"width_ratios": [1.3, 1, 1.4]}) + + ax = axes[0] + grid = [ + [0.5 if c.get(field) is None else float(c[field][item]) for field, item, _ in COLUMNS] + for c in cands + ] + ax.imshow(grid, cmap=ListedColormap([FAIL, NA, PASS]), vmin=0, vmax=1, aspect="auto") + ax.set_xticks(range(len(COLUMNS)), [c[2] for c in COLUMNS], fontsize=7, rotation=90) + ax.set_yticks(range(len(cands)), labels, fontsize=8) + for i, cand in enumerate(cands): + mark = "BI" if cand.get("batch_invariant") else "not BI" + ax.text(len(COLUMNS) - 0.4, i, f" {mark}", va="center", fontsize=8, color=INK) + ax.set_title( + f"Invariance: {report['target_tokens']}-token target alone vs. batched / chunked\n" + "(green bitwise equal, red differs, grey: no backward)", + fontsize=9, + color=INK, + ) + ax.tick_params(length=0) + + ax = axes[1] + _style(ax) + rel = [c["accuracy"]["rel_l2_vs_fp64"] for c in cands] + ax.barh(range(len(cands)), rel, color=[colors[k] for k in keys]) + for i, value in enumerate(rel): + ax.text(value, i, f" {value:.2e}", va="center", fontsize=7, color=MUTED) + ax.set_yticks(range(len(cands)), labels, fontsize=8) + ax.invert_yaxis() + ax.set_xlim(0, 1.35 * max(rel)) + ax.set_xlabel("relative L2 error vs. FP64", fontsize=8) + ax.set_title("Accuracy", fontsize=9, color=INK) + + ax = axes[2] + _style(ax) + prefill = report["latency"]["prefill_us"] + tokens = [int(t) for t in prefill] + for key in keys: + ys = [prefill[str(t)].get(key, float("nan")) / 1000.0 for t in tokens] + ax.plot(tokens, ys, marker="o", ms=3.5, color=colors[key], label=LABELS.get(key, key)) + ax.set_xscale("log", base=2) + ax.set_yscale("log") + ax.set_xlabel("prefill tokens", fontsize=8) + ax.set_ylabel("forward latency (ms, median)", fontsize=8) + ax.set_title("Prefill forward latency\n(4 Q heads, 1 KV head, D=256, causal)", fontsize=9) + ax.legend(fontsize=7, frameon=False) + + env = report["environment"] + fig.suptitle( + f"Qwen3-Next TP4 attention prior art - {env['gpu']}, torch {env['torch']}, " + f"commit {report['rl_kernel_commit'][:7]}", + fontsize=10, + color=INK, + ) + fig.tight_layout() + fig.savefig(out, dpi=160) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("report", type=Path) + args = parser.parse_args() + report = json.loads(args.report.read_text()) + out = args.report.with_name("figure.png") + plot(report, out) + print(f"wrote {out}") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/plot_qwen3_next_moe_prior_art.py b/tools/validation/models/plot_qwen3_next_moe_prior_art.py new file mode 100644 index 000000000..21cee4fb9 --- /dev/null +++ b/tools/validation/models/plot_qwen3_next_moe_prior_art.py @@ -0,0 +1,162 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Render the routed-MoE prior-art report as a PNG. + +Needs only the report JSON and matplotlib, so it runs anywhere: + + python tools/validation/models/plot_qwen3_next_moe_prior_art.py \\ + docs/usage/evidence/qwen3-next-moe-route-b200/report.json + +writes ``figure.png`` next to the report. Three panels: which batch-invariance +checks each implementation passes, accuracy against an FP64 evaluation with +FP64 routing, and forward latency. RL-Kernel is blue, existing libraries orange. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 +from matplotlib.colors import ListedColormap # noqa: E402 + +RL_KERNEL = "#2a78d6" +INK = "#2b2b29" +MUTED = "#6f6e69" +GRID = "#e4e3dd" +SURFACE = "#fcfcfb" +PASS = "#1baf7a" +FAIL = "#d64545" +NA = "#e4e3dd" +EXISTING = ("#eb6834", "#b8501f", "#f29a62", "#8a3a15", "#f6bd94", "#d9773f", "#6b2a0e") + +LABELS = { + "rl_kernel_cuda": "RL-Kernel shared_moe", + "hf_transformers": "HF Qwen3NextExperts", + "vllm_bi0": "vLLM fused_moe (BI=0)", + "vllm_bi1": "vLLM fused_moe (BI=1)", + "flashinfer_cutlass": "FlashInfer cutlass_fused_moe", + "megatron_te": "Megatron-core + TE (VIME config)", + "sglang_triton": "SGLang fused_moe (default)", + "sglang_deterministic": "SGLang fused_moe (deterministic)", +} +CHECKS = ( + ("route_rows_bitwise", "routes"), + ("output_rows_bitwise", "output"), + ("dx_rows_bitwise", "dx"), + ("dweight_zero_rows_bitwise", "dW"), +) + + +def _colors(keys): + existing = iter(EXISTING) + return {k: RL_KERNEL if k == "rl_kernel_cuda" else next(existing) for k in keys} + + +def _style(ax): + ax.set_facecolor(SURFACE) + for side in ("top", "right"): + ax.spines[side].set_visible(False) + ax.tick_params(colors=INK, labelsize=8) + ax.grid(color=GRID, linewidth=0.6) + + +def plot(report: dict, out: Path) -> None: + cands = [c for c in report["candidates"] if "unavailable" not in c] + keys = [c["key"] for c in cands] + colors = _colors(keys) + sizes = list(next(iter(cands))["output_rows_bitwise"]) + fig, axes = plt.subplots(1, 3, figsize=(17, 4.6), gridspec_kw={"width_ratios": [1.5, 1, 1.4]}) + fig.patch.set_facecolor("white") + + # Panel 1: batch-invariance matrix; 1 = pass, 0 = fail, 0.5 = no backward. + ax = axes[0] + columns = [f"{label} {size}" for _, label in CHECKS for size in sizes] + grid = [] + for cand in cands: + row = [] + for field, _ in CHECKS: + values = cand.get(field) + row.extend(0.5 if values is None else float(values[s]) for s in sizes) + grid.append(row) + ax.imshow(grid, cmap=ListedColormap([FAIL, NA, PASS]), vmin=0, vmax=1, aspect="auto") + ax.set_xticks(range(len(columns)), columns, fontsize=7, rotation=90) + ax.set_yticks(range(len(cands)), [LABELS.get(k, k) for k in keys], fontsize=8) + for i, cand in enumerate(cands): + mark = "BI" if cand.get("batch_invariant") else "not BI" + ax.text(len(columns) - 0.4, i, f" {mark}", va="center", fontsize=8, color=INK) + ax.set_title( + f"Batch invariance: {report['probe_rows']} probe tokens alone vs. in a batch\n" + "(green bitwise equal, red differs, grey: no backward)", + fontsize=9, + color=INK, + ) + ax.tick_params(length=0) + + # Panel 2: accuracy. + ax = axes[1] + _style(ax) + rel = [c["accuracy"]["rel_l2_vs_fp64"] if "accuracy" in c else float("nan") for c in cands] + ax.barh(range(len(cands)), rel, color=[colors[k] for k in keys]) + for i, cand in enumerate(cands): + acc = cand.get("accuracy") + if acc: + ax.text( + rel[i], + i, + f" {rel[i]:.2e} ({acc['tokens_with_different_expert_set']}/{acc['tokens']}" + " tokens route differently)", + va="center", + fontsize=7, + color=MUTED, + ) + ax.set_yticks(range(len(cands)), [LABELS.get(k, k) for k in keys], fontsize=8) + ax.invert_yaxis() + ax.set_xlim(0, 2.2 * max(r for r in rel if r == r)) + ax.set_xlabel("relative L2 error vs. FP64 HF formula with FP64 routing", fontsize=8) + ax.set_title("Accuracy (256 tokens)", fontsize=9, color=INK) + + # Panel 3: forward latency. + ax = axes[2] + _style(ax) + forward = report["latency"]["forward_us"] + tokens = [int(t) for t in forward] + for key in keys: + ys = [forward[str(t)].get(key, float("nan")) / 1000.0 for t in tokens] + ax.plot(tokens, ys, marker="o", ms=3.5, color=colors[key], label=LABELS.get(key, key)) + ax.set_xscale("log", base=2) + ax.set_yscale("log") + ax.set_xlabel("tokens", fontsize=8) + ax.set_ylabel("forward latency (ms, median)", fontsize=8) + ax.set_title("Forward latency\n(H=2048, 512 experts, top-10, width 512)", fontsize=9) + ax.legend(fontsize=7, frameon=False) + + env = report["environment"] + fig.suptitle( + f"Qwen3-Next routed MoE prior art - {env['gpu']}, torch {env['torch']}, " + f"commit {report['rl_kernel_commit'][:7]}", + fontsize=10, + color=INK, + ) + fig.tight_layout() + fig.savefig(out, dpi=160) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("report", type=Path) + args = parser.parse_args() + report = json.loads(args.report.read_text()) + out = args.report.with_name("figure.png") + plot(report, out) + print(f"wrote {out}") + + +if __name__ == "__main__": + main() 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_attention_prior_art.py b/tools/validation/models/qwen3_next_attention_prior_art.py new file mode 100644 index 000000000..6f4f0e485 --- /dev/null +++ b/tools/validation/models/qwen3_next_attention_prior_art.py @@ -0,0 +1,122 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Write the Qwen3-Next TP4 attention prior-art report (BI + accuracy + latency) as JSON. + +Run it on a clean tree with ``rl_engine._C`` built in place. The report records +the commit, whether the tracked tree was dirty and the version of every library +compared. +``tools/validation/models/plot_qwen3_next_attention_prior_art.py`` turns the report into a figure. + + python tools/validation/models/qwen3_next_attention_prior_art.py \\ + --out docs/usage/evidence/qwen3-next-tp4-mixers-b200/attention/report.json +""" + +from __future__ import annotations + +import argparse +import importlib.metadata as md +import json +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(ROOT)) + +import torch # noqa: E402 + +from rl_engine.validation.models.qwen3_next_attention_prior_art import ( # noqa: E402 + attention_report, +) + +LIBRARIES = ( + "vllm", + "flashinfer-python", + "transformers", + "triton", + "sglang", + "megatron-core", + "megatron_core", + "transformer_engine", + "vime", +) + + +def _git(*args: str) -> str: + return subprocess.run( + ["git", *args], cwd=ROOT, check=True, capture_output=True, text=True + ).stdout.strip() + + +def _version(name: str) -> str | None: + try: + return md.version(name) + except md.PackageNotFoundError: + return None + + +def _merge_latency(base: dict, extra: dict) -> dict: + """Per-candidate latencies merged at any depth (prefill is keyed by size first).""" + merged = dict(base) + for key, value in extra.items(): + if isinstance(value, dict) and isinstance(base.get(key), dict): + merged[key] = _merge_latency(base[key], value) + else: + merged[key] = value + return merged + + +def merge(base: dict, extra: dict) -> dict: + """``base`` with ``extra``'s candidates and latencies added (same commit and GPU).""" + for field in ("rl_kernel_commit", "tracked_tree_dirty", "shape", "op"): + if base[field] != extra[field]: + raise SystemExit(f"Reports differ in {field}; they cannot be merged") + if base["environment"]["gpu"] != extra["environment"]["gpu"]: + raise SystemExit("Reports were measured on different GPUs") + keys = {c["key"] for c in base["candidates"]} + added = [c for c in extra["candidates"] if c["key"] not in keys] + merged = dict(base, candidates=base["candidates"] + added) + merged["latency"] = _merge_latency(base["latency"], extra["latency"]) + merged["environments"] = {"main": base["environment"], "extra": extra["environment"]} + return merged + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--only", help="comma-separated candidate keys (default: all)") + parser.add_argument( + "--merge", type=Path, help="a report from another process to merge into this one" + ) + args = parser.parse_args() + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + report = { + "kind": "qwen3_next_prior_art_report", + "rfc": 428, + "rl_kernel_commit": _git("rev-parse", "HEAD"), + "tracked_tree_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": { + "gpu": torch.cuda.get_device_name(), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "libraries": {name: _version(name) for name in LIBRARIES}, + }, + **attention_report(args.only.split(",") if args.only else None), + } + if args.merge is not None: + report = merge(json.loads(args.merge.read_text()), report) + 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['rl_kernel_commit'][:7]}, " + f"dirty={report['tracked_tree_dirty']})" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/qwen3_next_gate_common.py b/tools/validation/models/qwen3_next_gate_common.py new file mode 100644 index 000000000..3de49bdc1 --- /dev/null +++ b/tools/validation/models/qwen3_next_gate_common.py @@ -0,0 +1,125 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Shared pieces of the Qwen3-Next TP4 real-weight gates (``scripts/qwen3_next_tp_*_check.py``). + +Each gate is launched as ``torchrun --nproc-per-node 4 scripts/.py +--checkpoint --output `` and writes one ``rank-.json``. +""" + +import argparse +import json +import os +import subprocess +from collections.abc import Mapping +from datetime import timedelta +from importlib.metadata import version +from pathlib import Path + +import torch +import torch.distributed as dist +from safetensors import safe_open + +ROOT = Path(__file__).resolve().parents[3] + + +class CheckpointWeights(Mapping): + """Keep shard headers open while reading only the requested layer's tensors.""" + + def __init__(self, checkpoint, prefix, names, stack): + self.root, self.prefix, self.names, self.stack = checkpoint, prefix, tuple(names), stack + self.index = json.loads((checkpoint / "model.safetensors.index.json").read_text())[ + "weight_map" + ] + self.readers = {} + + def __iter__(self): + return iter(self.names) + + def __len__(self): + return len(self.names) + + def __getitem__(self, name): + if name not in self.names: + raise KeyError(name) + key = self.prefix + name + shard = self.index[key] + if shard not in self.readers: + self.readers[shard] = self.stack.enter_context( + safe_open(self.root / shard, framework="pt", device="cpu") + ) + return self.readers[shard].get_tensor(key) + + +def gather(value): + peers = [torch.empty_like(value) for _ in range(4)] + dist.all_gather(peers, value.contiguous()) + return peers + + +def check_gradients(module): + result = {} + for name, parameter in module.named_parameters(): + grad = parameter.grad + if grad is None or not torch.isfinite(grad).all() or not grad.count_nonzero(): + raise AssertionError(f"Missing finite nonzero gradient: {name}") + result[name] = float(grad.float().abs().max()) + return result + + +def timed(fn): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + out = fn() + end.record() + torch.cuda.synchronize() + return out, start.elapsed_time(end) * 1000.0 + + +def _git(*args): + return subprocess.run( + ["git", *args], cwd=ROOT, check=True, capture_output=True, text=True + ).stdout.strip() + + +def provenance(): + from rl_engine.models.qwen3_next.qwen3_next_forward import FORWARD_PROVIDER_ID + + return { + "rl_kernel_commit": _git("rev-parse", "HEAD"), + "tracked_tree_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "provider": FORWARD_PROVIDER_ID, + "environment": { + "gpu": torch.cuda.get_device_name(), + "torch": torch.__version__, + "vllm": version("vllm"), + "nccl": ".".join(map(str, torch.cuda.nccl.version())), + }, + } + + +def run_gate(description, scope, gate): + """Parse ``--checkpoint/--output``, run ``gate(checkpoint, generator)`` on four ranks.""" + parser = argparse.ArgumentParser(description=description) + parser.add_argument("--checkpoint", required=True, type=Path) + parser.add_argument("--output", required=True, type=Path) + args = parser.parse_args() + rank = int(os.environ["LOCAL_RANK"]) + if int(os.environ["WORLD_SIZE"]) != 4 or torch.cuda.device_count() < 4: + raise RuntimeError(f"The {scope} gate requires exactly four ranks on four GPUs") + torch.cuda.set_device(rank) + dist.init_process_group( + "nccl", timeout=timedelta(minutes=15), device_id=torch.device("cuda", rank) + ) + try: + generator = torch.Generator(device="cuda").manual_seed(1234) + result = gate(args.checkpoint, generator) + args.output.mkdir(parents=True, exist_ok=True) + path = args.output / f"rank-{rank}.json" + partial = path.with_suffix(".json.partial") + payload = {"status": "passed", "scope": scope, "rank": rank, **provenance(), **result} + partial.write_text(json.dumps(payload, indent=2) + "\n") + partial.replace(path) + dist.barrier() + finally: + dist.destroy_process_group() diff --git a/tools/validation/models/qwen3_next_moe_prior_art.py b/tools/validation/models/qwen3_next_moe_prior_art.py new file mode 100644 index 000000000..c3c8d088a --- /dev/null +++ b/tools/validation/models/qwen3_next_moe_prior_art.py @@ -0,0 +1,116 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Write the Qwen3-Next routed-MoE prior-art report (BI + accuracy + latency) as JSON. + +Run it on a clean tree with ``TRITON_F32_DEFAULT=ieee`` (the FP32 router needs +IEEE Triton dots from process start). The report records the commit, whether +the tracked tree was dirty and the version of every library compared. +``tools/validation/models/plot_qwen3_next_moe_prior_art.py`` turns the report into a figure. + + TRITON_F32_DEFAULT=ieee python tools/validation/models/qwen3_next_moe_prior_art.py \\ + --out docs/usage/evidence/qwen3-next-moe-route-b200/report.json +""" + +from __future__ import annotations + +import argparse +import importlib.metadata as md +import json +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(ROOT)) + +import torch # noqa: E402 + +from rl_engine.validation.models.qwen3_next_moe_prior_art import moe_report # noqa: E402 + +LIBRARIES = ( + "vllm", + "flashinfer-python", + "transformers", + "triton", + "sglang", + "sgl-kernel", + "megatron_core", + "megatron-core", + "transformer_engine", + "vime", +) + + +def _git(*args: str) -> str: + return subprocess.run( + ["git", *args], cwd=ROOT, check=True, capture_output=True, text=True + ).stdout.strip() + + +def _version(name: str) -> str | None: + try: + return md.version(name) + except md.PackageNotFoundError: + return None + + +def merge(base: dict, extra: dict) -> dict: + """``base`` with ``extra``'s candidates and latencies added (same commit and GPU).""" + for field in ("rl_kernel_commit", "tracked_tree_dirty", "shape", "op"): + if base[field] != extra[field]: + raise SystemExit(f"Reports differ in {field}; they cannot be merged") + if base["environment"]["gpu"] != extra["environment"]["gpu"]: + raise SystemExit("Reports were measured on different GPUs") + keys = {c["key"] for c in base["candidates"]} + added = [c for c in extra["candidates"] if c["key"] not in keys] + merged = dict(base, candidates=base["candidates"] + added) + merged["latency"] = { + section: { + size: {**values, **extra["latency"][section].get(size, {})} + for size, values in sizes.items() + } + for section, sizes in base["latency"].items() + } + merged["environments"] = {"main": base["environment"], "extra": extra["environment"]} + return merged + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--only", help="comma-separated candidate keys (default: all)") + parser.add_argument( + "--merge", type=Path, help="a report from another environment to merge into this one" + ) + args = parser.parse_args() + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + report = { + "kind": "qwen3_next_prior_art_report", + "rfc": 428, + "rl_kernel_commit": _git("rev-parse", "HEAD"), + "tracked_tree_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": { + "gpu": torch.cuda.get_device_name(), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "libraries": {name: _version(name) for name in LIBRARIES}, + }, + **moe_report(args.only.split(",") if args.only else None), + } + if args.merge is not None: + report = merge(json.loads(args.merge.read_text()), report) + 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['rl_kernel_commit'][:7]}, " + f"dirty={report['tracked_tree_dirty']})" + ) + + +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/qwen3_next_tp_attention_check.py b/tools/validation/models/qwen3_next_tp_attention_check.py new file mode 100644 index 000000000..a94487e35 --- /dev/null +++ b/tools/validation/models/qwen3_next_tp_attention_check.py @@ -0,0 +1,145 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""TP4 full-attention gate on the official checkpoint's layer-3 weights; four GPUs, torchrun. + + torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_attention_check.py \\ + --checkpoint --output + +Requires, bitwise: the HF load/export round trip with the KV head replicated on +each rank pair; full prefill == chunked prefill (output and KV continuation) for +8-1024 tokens; a variable-length batch of four == each sequence alone, in any +order; a sixteen-token decode == the prefill rows; the training recompute == the +no-grad forward; and identical KV-pair and QK-norm gradients across replicas. +""" + +import sys +from contextlib import ExitStack +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 +import torch.distributed as dist # noqa: E402 + +from rl_engine.models.qwen3_next.qwen3_next_tp_blocks import ( # noqa: E402 + ATTENTION_HF_SHAPES, + TP4FullAttention, + assemble_attention_weights, +) +from rl_engine.validation.common.tensor_identity import ( # noqa: E402 + assert_tensor_bitwise_equal as exact, +) +from tools.validation.models.qwen3_next_gate_common import ( # noqa: E402 + CheckpointWeights, + check_gradients, + gather, + run_gate, +) + + +def attention_gate(checkpoint, generator): + module = TP4FullAttention(group=dist.group.WORLD, device="cuda") + with ExitStack() as stack: + weights = CheckpointWeights( + checkpoint, "model.layers.3.self_attn.", ATTENTION_HF_SHAPES, stack + ) + module.load_hf_weights(weights) + shards = [dict() for _ in range(4)] + for name, tensor in module.export_local_hf_weights().items(): + for rank, value in enumerate(gather(tensor)): + shards[rank][name] = value.cpu() + for name, tensor in assemble_attention_weights(shards).items(): + exact(tensor, weights[name], name=f"attention HF {name}") + del shards, weights + cases = ["hf-load-export"] + slots = torch.tensor([1], device="cuda", dtype=torch.int32) + with torch.no_grad(): + for length in (8, 64, 256, 1024): + hidden = ( + torch.randn(length, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) + * 0.1 + ) + state = module.initial_state(2) + whole, expected = module(hidden, state, slots, torch.tensor([0, length])) + cut = min(65, length - 1) + first, state = module(hidden[:cut], state, slots, torch.tensor([0, cut])) + last, state = module(hidden[cut:], state, slots, torch.tensor([0, length - cut])) + exact(torch.cat((first, last)), whole, name=f"attention full/chunk {length}") + exact(state.keys[1], expected.keys[1], name="attention key continuation") + exact(state.values[1], expected.values[1], name="attention value continuation") + cases.append(f"full-chunk-{length}") + hidden = ( + torch.randn(88, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) * 0.1 + ) + packed, state = module( + hidden, + module.initial_state(5), + torch.tensor([1, 2, 3, 4], device="cuda"), + torch.tensor([0, 8, 24, 48, 88]), + ) + order = [3, 1, 0, 2] + segments = [hidden[:8], hidden[8:24], hidden[24:48], hidden[48:]] + ends = [0] + for index in order: + ends.append(ends[-1] + segments[index].shape[0]) + reordered, _ = module( + torch.cat([segments[i] for i in order]), + module.initial_state(5), + torch.tensor([i + 1 for i in order], device="cuda"), + torch.tensor(ends), + ) + expected = packed.split((8, 16, 24, 40)) + exact( + reordered, + torch.cat([expected[i] for i in order]), + name="attention variable-length batch4 reorder", + ) + for index, segment in enumerate(segments): + single, _ = module( + segment, module.initial_state(2), slots, torch.tensor([0, segment.shape[0]]) + ) + exact(single, expected[index], name="attention batch4/single") + cases.append("batch4-varlen-reorder") + hidden = ( + torch.randn(80, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) * 0.1 + ).requires_grad_() + with torch.no_grad(): + expected, _ = module(hidden, module.initial_state(2), slots, torch.tensor([0, 80])) + _, decode_state = module(hidden[:64], module.initial_state(2), slots, torch.tensor([0, 64])) + decoded = [] + for token in hidden[64:].split(1): + value, decode_state = module(token, decode_state, slots, torch.tensor([0, 1])) + decoded.append(value) + exact(torch.cat(decoded), expected[64:], name="attention sixteen-token decode") + _, state = module(hidden[:64], module.initial_state(2), slots, torch.tensor([0, 64])) + output, _ = module(hidden[64:], state, slots, torch.tensor([0, 16])) + exact(output, expected[64:], name="attention independent response recompute") + output.float().square().mean().backward() + if ( + hidden.grad is None + or not torch.isfinite(hidden.grad).all() + or not hidden.grad[:64].count_nonzero() + ): + raise AssertionError("Attention response gradient did not reach the prompt") + gradients = check_gradients(module) + for projection in (module.k_proj, module.v_proj): + peers = gather(projection.weight.grad) + exact(peers[0], peers[1], name="KV replica gradient pair0") + exact(peers[2], peers[3], name="KV replica gradient pair1") + for norm in (module.q_norm, module.k_norm): + for peer in gather(norm.weight.grad): + exact(peer, norm.weight.grad, name="QK norm replicated gradient") + cases.extend( + ("sixteen-token-decode", "prompt-gradient", "kv-pair-gradient", "qk-norm-gradient") + ) + return {"cases": cases, "gradient_max_abs": gradients} + + +if __name__ == "__main__": + run_gate( + __doc__, + "tp4_full_attention_layer3", + lambda checkpoint, generator: {"attention": attention_gate(checkpoint, generator)}, + ) diff --git a/tools/validation/models/qwen3_next_tp_gdn_check.py b/tools/validation/models/qwen3_next_tp_gdn_check.py new file mode 100644 index 000000000..65fee3815 --- /dev/null +++ b/tools/validation/models/qwen3_next_tp_gdn_check.py @@ -0,0 +1,108 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""TP4 GDN gate on the official checkpoint's layer-0 weights; four GPUs, torchrun. + + torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_gdn_check.py \\ + --checkpoint --output + +A one-layer integration gate, not full-model evidence. Requires, bitwise: the +HF load/export round trip of the actual parameter storage; full prefill == +chunked prefill, output and convolution/recurrent state handoff, for 8-1024 +tokens; the training response recompute == the no-grad forward; a gradient that +reaches the prompt through the state; identical replicated norm gradients. +""" + +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 +import torch.distributed as dist # noqa: E402 +from safetensors import safe_open # noqa: E402 + +from rl_engine.integrations.engines.train.vllm.qwen3_next_provider import ( # noqa: E402 + GDNProviderConfig, +) +from rl_engine.models.qwen3_next.qwen3_next_tp_gdn import ( # noqa: E402 + _GLOBAL_SHAPES, + TP4GDN, + assemble_gdn_weights, +) +from rl_engine.validation.common.tensor_identity import ( # noqa: E402 + assert_tensor_bitwise_equal as exact, +) +from tools.validation.models.qwen3_next_gate_common import check_gradients, run_gate # noqa: E402 + + +def load_first_gdn(checkpoint): + index = json.loads((checkpoint / "model.safetensors.index.json").read_text())["weight_map"] + result = {} + for name in _GLOBAL_SHAPES: + key = f"model.layers.0.linear_attn.{name}" + with safe_open(checkpoint / index[key], framework="pt", device="cpu") as reader: + result[name] = reader.get_tensor(key) + return result + + +def gdn_gate(checkpoint, generator): + weights = load_first_gdn(checkpoint) + module = TP4GDN(group=dist.group.WORLD, device="cuda", config=GDNProviderConfig()) + module.load_hf_weights(weights) + # Validate actual parameter storage, not just the pure conversion helper. + shards = [dict() for _ in range(4)] + for name, parameter in module.named_parameters(): + gathered = [torch.empty_like(parameter) for _ in range(4)] + dist.all_gather(gathered, parameter.detach()) + for peer, value in enumerate(gathered): + shards[peer][name] = value.cpu() + for name, restored in assemble_gdn_weights(shards).items(): + exact(restored, weights[name], name=f"loaded weight {name}") + del shards, weights + cases = [] + indices = torch.tensor([1], device="cuda", dtype=torch.int32) + with torch.no_grad(): + for length in (8, 64, 256, 1024): + hidden = ( + torch.randn(length, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) + * 0.1 + ) + initial = module.initial_state(2) + whole, expected_state = module(hidden, initial, indices, torch.tensor([0, length])) + cut = min(65, length - 1) + first, state = module(hidden[:cut], initial, indices, torch.tensor([0, cut])) + last, state = module(hidden[cut:], state, indices, torch.tensor([0, length - cut])) + exact(torch.cat((first, last)), whole, name=f"TP4 layer length={length}") + exact(state.convolution, expected_state.convolution, name="convolution handoff") + exact(state.recurrent, expected_state.recurrent, name="recurrent handoff") + cases.append(f"full-chunk-{length}") + hidden = ( + torch.randn(80, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) * 0.1 + ).requires_grad_() + initial = module.initial_state(2) + with torch.no_grad(): + expected, _ = module(hidden, initial, indices, torch.tensor([0, 80])) + _, prompt_state = module(hidden[:64], initial, indices, torch.tensor([0, 64])) + actual, _ = module(hidden[64:], prompt_state, indices, torch.tensor([0, 16])) + exact(actual, expected[64:], name="independent training response") + actual.float().square().mean().backward() + if ( + hidden.grad is None + or not torch.isfinite(hidden.grad).all() + or not hidden.grad[:64].count_nonzero() + ): + raise AssertionError("Response gradient did not reach prompt input") + grad_stats = check_gradients(module) + norm_grads = [torch.empty_like(module.norm.weight.grad) for _ in range(4)] + dist.all_gather(norm_grads, module.norm.weight.grad) + for other in norm_grads: + exact(other, module.norm.weight.grad, name="replicated norm gradient") + cases.extend(("prompt-gradient", "replicated-norm-gradient", "hf-load-export")) + return {"cases": cases, "gradient_max_abs": grad_stats} + + +if __name__ == "__main__": + run_gate(__doc__, "tp4_gdn_layer0", gdn_gate) diff --git a/tools/validation/models/qwen3_next_tp_moe_check.py b/tools/validation/models/qwen3_next_tp_moe_check.py new file mode 100644 index 000000000..c18b5b1fd --- /dev/null +++ b/tools/validation/models/qwen3_next_tp_moe_check.py @@ -0,0 +1,114 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""TP4 MoE gate on the official checkpoint's layer-0 weights; four GPUs, torchrun. + + TRITON_F32_DEFAULT=ieee torchrun --nproc-per-node 4 \\ + tools/validation/models/qwen3_next_tp_moe_check.py \\ + --checkpoint --output + +Each rank writes ``rank-.json``. The gate requires, bitwise: + +* the HF load -> local shard -> all-gather -> HF export round trip of all 512 experts; +* identical top-10 routes on every rank (the router is replicated); +* full batch == concatenated chunks, and reordered tokens == reordered output, + for 8, 64, 256 and 1024 tokens; +* the training forward (autograd on) == the no-grad forward; +* identical replicated router and shared-expert-gate gradients on every rank, + and a finite, nonzero gradient for every parameter. +""" + +import sys +from contextlib import ExitStack +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 +import torch.distributed as dist # noqa: E402 + +from rl_engine.models.qwen3_next.qwen3_next_forward import ( # noqa: E402 + shared_router, + stable_top10_routes, +) +from rl_engine.models.qwen3_next.qwen3_next_tp_blocks import ( # noqa: E402 + TP4MoE, + assemble_moe_weight, + moe_hf_shapes, +) +from rl_engine.validation.common.tensor_identity import ( # noqa: E402 + assert_tensor_bitwise_equal as exact, +) +from tools.validation.models.qwen3_next_gate_common import ( # noqa: E402 + CheckpointWeights, + check_gradients, + gather, + run_gate, + timed, +) + +TOKEN_COUNTS = (8, 64, 256, 1024) + + +def moe_gate(checkpoint, generator): + module = TP4MoE(group=dist.group.WORLD, device="cuda") + with ExitStack() as stack: + weights = CheckpointWeights(checkpoint, "model.layers.0.mlp.", moe_hf_shapes(), stack) + module.load_hf_weights(weights) + for name, value in module.export_local_hf_weights().items(): + restored = assemble_moe_weight(name, [peer.cpu() for peer in gather(value)]) + exact(restored, weights[name], name=f"MoE HF {name}") + cases = ["hf-load-export-all512experts"] + forward_us = {} + with torch.no_grad(): + for length in TOKEN_COUNTS: + hidden = ( + torch.randn(length, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) + * 0.1 + ) + # Every rank draws the same tokens from the same generator; verify it. + for peer in gather(hidden): + exact(peer, hidden, name="replicated input") + routes = stable_top10_routes(shared_router(hidden, module.gate.weight)) + for peer in gather(routes.indices): + exact(peer, routes.indices, name=f"routes across ranks {length}") + module(hidden) + whole, forward_us[str(length)] = timed(lambda: module(hidden)) + cut = min(65, length - 1) + exact( + torch.cat((module(hidden[:cut]), module(hidden[cut:]))), + whole, + name=f"MoE full/chunk {length}", + ) + order = torch.arange(length - 1, -1, -1, device="cuda") + exact(module(hidden[order]), whole[order], name=f"MoE reorder {length}") + cases.extend((f"routes-replicated-{length}", f"full-chunk-reorder-{length}")) + hidden = ( + torch.randn(16, 2048, device="cuda", dtype=torch.bfloat16, generator=generator) * 0.1 + ).requires_grad_() + with torch.no_grad(): + expected = module(hidden) + output = module(hidden) + exact(output, expected, name="MoE independent training forward") + output.float().square().mean().backward() + if ( + hidden.grad is None + or not torch.isfinite(hidden.grad).all() + or not hidden.grad.count_nonzero() + ): + raise AssertionError("MoE input gradient is missing") + gradients = check_gradients(module) + for parameter in (module.gate.weight, module.shared_expert_gate.weight): + for peer in gather(parameter.grad): + exact(peer, parameter.grad, name="MoE replicated router/gate gradient") + cases.extend(("training-forward", "all-parameter-gradient", "replicated-router-gradient")) + return {"cases": cases, "gradient_max_abs": gradients, "forward_us": forward_us} + + +if __name__ == "__main__": + run_gate( + __doc__, + "tp4_moe_block_layer0", + lambda checkpoint, generator: {"moe": moe_gate(checkpoint, generator)}, + ) 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: