diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4634569cd..a73ecfa02 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,9 +5,9 @@ name: CI-Pipeline on: push: - branches: [ main, test ] + branches: [ main, test, test-qwennext ] pull_request: - branches: [ main, test ] + branches: [ main, test, test-qwennext ] permissions: contents: read @@ -75,7 +75,7 @@ jobs: run: | python -m pytest tests/runtime/test_dispatch.py -v PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest tests/ops/attention/test_attention_correctness.py -q -rs - python -m pytest tests/validation/operators/test_forward_invariance.py tests/contracts/test_tolerance_contract.py tests/validation/operators/test_ws1_workload.py tests/validation/operators/test_gradient_invariance.py tests/validation/operators/test_elementwise_inventory.py tests/validation/operators/test_four_judgment_matrix.py tests/validation/operators/test_op_checks.py tests/validation/operators/test_operator_inputs.py tests/benchmarks/test_profiler.py tests/validation/operators/test_kv_consistency.py tests/models/qwen3/test_ws1_qwen3_dense.py tests/validation/models/test_ws1_chain_integration.py -q + python -m pytest tests/validation/operators/test_forward_invariance.py tests/contracts/test_tolerance_contract.py tests/validation/operators/test_ws1_workload.py tests/validation/operators/test_gradient_invariance.py tests/validation/operators/test_elementwise_inventory.py tests/validation/operators/test_four_judgment_matrix.py tests/validation/operators/test_op_checks.py tests/validation/operators/test_operator_inputs.py tests/benchmarks/test_profiler.py tests/validation/operators/test_kv_consistency.py tests/models/qwen3/test_ws1_qwen3_dense.py tests/validation/models/test_ws1_chain_integration.py tests/models/qwen3_next/test_qwen3_next_norm.py -q - name: Run Cross-Configuration Contract Tests (CPU-safe) run: | diff --git a/.github/workflows/ws1-gtest-gpu.yml b/.github/workflows/ws1-gtest-gpu.yml index a0d8ac0c4..b853b9df0 100644 --- a/.github/workflows/ws1-gtest-gpu.yml +++ b/.github/workflows/ws1-gtest-gpu.yml @@ -10,7 +10,7 @@ name: WS1-gtest-GPU on: pull_request: - branches: [ main ] + branches: [ main, test-qwennext ] paths: - "rl_engine/validation/**" - "rl_engine/backends/**" @@ -28,6 +28,7 @@ on: - "tools/validation/operators/ws1_candidate_evidence.py" - "tests/validation/**" - "tests/models/**" + - "tests/ops/norm/**" - "tests/validation/operators/test_forward_invariance.py" - "tests/validation/operators/test_gradient_invariance.py" - "tests/validation/operators/test_four_judgment_matrix.py" @@ -36,7 +37,7 @@ on: - "ci/scripts/run_gpu_ci.sh" - ".github/workflows/ws1-gtest-gpu.yml" push: - branches: [ main ] + branches: [ main, test-qwennext ] paths: - "rl_engine/validation/**" - "rl_engine/backends/**" @@ -54,6 +55,7 @@ on: - "tools/validation/operators/ws1_candidate_evidence.py" - "tests/validation/**" - "tests/models/**" + - "tests/ops/norm/**" - "tests/validation/operators/test_forward_invariance.py" - "tests/validation/operators/test_gradient_invariance.py" - "tests/validation/operators/test_four_judgment_matrix.py" diff --git a/ci/scripts/run_ws1_gtest.sh b/ci/scripts/run_ws1_gtest.sh index cc47dd678..1c5a2be53 100755 --- a/ci/scripts/run_ws1_gtest.sh +++ b/ci/scripts/run_ws1_gtest.sh @@ -19,6 +19,7 @@ echo "[ws1-gtest] interpreter=$PY out=$OUT" "$PY" -m pytest -q \ tests/validation/operators/test_ws1_gtest_gpu.py \ + tests/models/qwen3_next/test_qwen3_next_norm.py \ tests/ops/attention/test_triton_batch_invariant_attention.py \ tests/validation/operators/test_four_judgment_matrix.py \ tests/validation/operators/test_ws1_candidate_evidence.py \ @@ -75,3 +76,14 @@ for cell in required: ) print("[ws1-gtest] C8 gate passed") PY + +echo "[ws1-gtest] Qwen3-Next C3/C4 norms (operator scope)" +QWEN3_NEXT_NORM_MANIFEST=rl_engine/validation/models/qwen3_next_norm_manifest.json +"$PY" tools/validation/operators/check_forward_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op qwen3_next_rms_norm --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_gradient_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op qwen3_next_rms_norm --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_forward_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op rms_norm_gated --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 +"$PY" tools/validation/operators/check_gradient_invariance.py --manifest "$QWEN3_NEXT_NORM_MANIFEST" \ + --op rms_norm_gated --candidate cuda --backend-profile cuda_bf16 --hidden 2048 --head-dim 128 diff --git a/csrc/bindings/ops.cpp b/csrc/bindings/ops.cpp index 0c7c3e7f4..a6e0d52f0 100644 --- a/csrc/bindings/ops.cpp +++ b/csrc/bindings/ops.cpp @@ -263,14 +263,36 @@ void rmsnorm_forward_cuda( torch::Tensor weight, torch::Tensor y, torch::Tensor rstd, - double eps); + double eps, + double weight_offset); void rmsnorm_backward_dx_cuda( torch::Tensor dy, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, - torch::Tensor dx); + torch::Tensor dx, + double weight_offset); + +void rmsnorm_gated_forward_cuda( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor y, + torch::Tensor rstd, + double eps, + double weight_offset, + int64_t activation); + +void rmsnorm_gated_backward_dx_cuda( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + torch::Tensor dx, + double weight_offset, + int64_t activation); void rmsnorm_backward_partial_dw_cuda( torch::Tensor dy, @@ -296,23 +318,39 @@ static void rmsnorm_check_input(const torch::Tensor& x, const char* name) { TORCH_CHECK(x.is_contiguous(), name, " must be contiguous"); } +static void rmsnorm_check_weight(const torch::Tensor& x, const torch::Tensor& weight) { + TORCH_CHECK(x.dim() == 2 && x.size(1) > 0, "x must be 2D with positive hidden size"); + TORCH_CHECK(x.scalar_type() == torch::kFloat32 || x.scalar_type() == torch::kFloat16 || + x.scalar_type() == torch::kBFloat16, "x must be float32, float16 or bfloat16"); + TORCH_CHECK(weight.dim() == 1 && weight.size(0) == x.size(1), "weight must be [H]"); + TORCH_CHECK(weight.device() == x.device(), "weight must be on the same device as x"); + TORCH_CHECK(weight.scalar_type() == x.scalar_type() || weight.scalar_type() == torch::kFloat32, + "weight must have x dtype or float32"); +} + +static void rmsnorm_check_backward(const torch::Tensor& dy, const torch::Tensor& x, + const torch::Tensor& rstd) { + TORCH_CHECK(dy.device() == x.device() && rstd.device() == x.device(), + "dy and rstd must be on the same device as x"); + TORCH_CHECK(dy.scalar_type() == x.scalar_type(), "dy must have the same dtype as x"); + TORCH_CHECK(rstd.scalar_type() == torch::kFloat32, "rstd must be float32"); +} + std::vector rmsnorm_forward( torch::Tensor x, torch::Tensor weight, - double eps) + double eps, + double weight_offset) { rmsnorm_check_input(x, "x"); rmsnorm_check_input(weight, "weight"); - - TORCH_CHECK(x.dim() == 2, "x must be 2D [T, H]"); - TORCH_CHECK(weight.dim() == 1, "weight must be 1D [H]"); - TORCH_CHECK(x.size(1) == weight.size(0), "x.size(1) must equal weight.size(0)"); + rmsnorm_check_weight(x, weight); auto T = x.size(0); auto y = torch::empty_like(x); auto rstd = torch::empty({T}, x.options().dtype(torch::kFloat32)); - rmsnorm_forward_cuda(x, weight, y, rstd, eps); + rmsnorm_forward_cuda(x, weight, y, rstd, eps, weight_offset); return {y, rstd}; } @@ -321,22 +359,89 @@ torch::Tensor rmsnorm_backward_dx( torch::Tensor dy, torch::Tensor x, torch::Tensor weight, - torch::Tensor rstd) + torch::Tensor rstd, + double weight_offset) { rmsnorm_check_input(dy, "dy"); rmsnorm_check_input(x, "x"); rmsnorm_check_input(weight, "weight"); + rmsnorm_check_weight(x, weight); rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); TORCH_CHECK(dy.sizes() == x.sizes(), "dy and x must have same shape"); - TORCH_CHECK(x.dim() == 2, "x must be 2D [T, H]"); - TORCH_CHECK(weight.dim() == 1, "weight must be 1D [H]"); TORCH_CHECK(rstd.dim() == 1, "rstd must be 1D [T]"); TORCH_CHECK(rstd.size(0) == x.size(0), "rstd.size(0) must equal x.size(0)"); auto dx = torch::empty_like(x); - rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx); + rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx, weight_offset); + + return dx; +} + +static void rmsnorm_gated_check( + const torch::Tensor& x, + const torch::Tensor& weight, + const torch::Tensor& gate, + int64_t activation) +{ + rmsnorm_check_input(x, "x"); + rmsnorm_check_input(weight, "weight"); + rmsnorm_check_weight(x, weight); + rmsnorm_check_input(gate, "gate"); + TORCH_CHECK(gate.device() == x.device(), "gate must be on the same device as x"); + TORCH_CHECK(gate.sizes() == x.sizes(), "gate must have the same shape as x"); + TORCH_CHECK( + gate.scalar_type() == x.scalar_type(), + "gate must have the same dtype as x"); + // 0 = silu/swish, 1 = sigmoid. Anything else fails closed rather than + // silently computing a different activation (RFC #428 section 6, item 7). + TORCH_CHECK( + activation == 0 || activation == 1, + "activation must be 0 (silu) or 1 (sigmoid), got ", activation); +} + +std::vector rmsnorm_gated_forward( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + double eps, + double weight_offset, + int64_t activation) +{ + rmsnorm_gated_check(x, weight, gate, activation); + + auto T = x.size(0); + auto y = torch::empty_like(x); + auto rstd = torch::empty({T}, x.options().dtype(torch::kFloat32)); + + rmsnorm_gated_forward_cuda(x, weight, gate, y, rstd, eps, weight_offset, activation); + + return {y, rstd}; +} + +torch::Tensor rmsnorm_gated_backward_dx( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + double weight_offset, + int64_t activation) +{ + rmsnorm_gated_check(x, weight, gate, activation); + rmsnorm_check_input(dy, "dy"); + rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); + + TORCH_CHECK(dy.sizes() == x.sizes(), "dy must have the same shape as x"); + TORCH_CHECK(rstd.dim() == 1 && rstd.size(0) == x.size(0), "rstd must be [T]"); + + auto dx = torch::empty_like(x); + + rmsnorm_gated_backward_dx_cuda( + dy, x, weight, gate, rstd, dx, weight_offset, activation); return dx; } @@ -350,6 +455,7 @@ torch::Tensor rmsnorm_backward_dw( rmsnorm_check_input(dy, "dy"); rmsnorm_check_input(x, "x"); rmsnorm_check_input(rstd, "rstd"); + rmsnorm_check_backward(dy, x, rstd); rmsnorm_check_input(mask, "mask"); TORCH_CHECK(dy.sizes() == x.sizes(), "dy and x must have same shape"); @@ -357,6 +463,8 @@ torch::Tensor rmsnorm_backward_dw( TORCH_CHECK(rstd.dim() == 1, "rstd must be 1D [T]"); TORCH_CHECK(mask.dim() == 1, "mask must be 1D [T]"); TORCH_CHECK(mask.scalar_type() == torch::kBool, "mask must be bool"); + TORCH_CHECK(mask.device() == x.device(), "mask must be on the same device as x"); + TORCH_CHECK(x.size(1) > 0, "hidden size must be positive"); TORCH_CHECK(rstd.size(0) == x.size(0), "rstd.size(0) must equal x.size(0)"); TORCH_CHECK(mask.size(0) == x.size(0), "mask.size(0) must equal x.size(0)"); @@ -718,9 +826,24 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &det_gemm_db_transposed, "Batch-invariant deterministic GEMM backward in canonical [N,K] layout"); // registry RMSNorm - m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA"); - m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA"); + // API version 2 includes weight_offset in both entry points. + m.attr("rmsnorm_api_version") = 2; + m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA", + py::arg("x"), py::arg("weight"), py::arg("eps"), + py::arg("weight_offset") = 0.0); + m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA", + py::arg("dy"), py::arg("x"), py::arg("weight"), py::arg("rstd"), + py::arg("weight_offset") = 0.0); m.def("rmsnorm_backward_dw", &rmsnorm_backward_dw, "Deterministic RMSNorm backward dweight CUDA"); + m.def("rmsnorm_gated_forward", &rmsnorm_gated_forward, + "Batch-invariant gated RMSNorm forward CUDA", + py::arg("x"), py::arg("weight"), py::arg("gate"), py::arg("eps"), + py::arg("weight_offset") = 0.0, py::arg("activation") = 0); + m.def("rmsnorm_gated_backward_dx", &rmsnorm_gated_backward_dx, + "Batch-invariant gated RMSNorm backward dx CUDA", + py::arg("dy"), py::arg("x"), py::arg("weight"), py::arg("gate"), + py::arg("rstd"), py::arg("weight_offset") = 0.0, + py::arg("activation") = 0); #if !defined(USE_ROCM) m.def( "reduce_rows_fp32_left_fold", diff --git a/csrc/cuda/norm/rmsnorm.cu b/csrc/cuda/norm/rmsnorm.cu index b32bc5af4..81a02cd2d 100644 --- a/csrc/cuda/norm/rmsnorm.cu +++ b/csrc/cuda/norm/rmsnorm.cu @@ -108,7 +108,8 @@ __global__ void rmsnorm_fwd_kernel( float* __restrict__ rstd, int T, int H, - float eps + float eps, + float weight_offset ) { int row = blockIdx.x; int tid = threadIdx.x; @@ -135,10 +136,17 @@ __global__ void rmsnorm_fwd_kernel( __syncthreads(); - // Write y = x * rstd * weight. + // Write y = x * rstd * (weight_offset + weight). The offset is added in + // fp32 after the upcast: a zero-centred weight (Qwen3-Next, Gemma) must not + // have its "+1" folded into the low-precision weight beforehand, which would + // round the offset and break the bitwise contract. for (int col = tid; col < H; col += blockDim.x) { float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + // Guarded: `-0.0f + 0.0f` is +0.0f, so an unconditional add would flip the + // sign bit of -0.0 weights on the plain path. Pinned by + // tests/models/qwen3_next/test_qwen3_next_norm.py::test_cuda_plain_signed_zero_is_preserved. + if (weight_offset != 0.0f) wv += weight_offset; float out = xv * row_rstd * wv; store_from_float(y_row + col, out); } @@ -153,7 +161,8 @@ __global__ void rmsnorm_bwd_dx_kernel( const float* __restrict__ rstd, scalar_t* __restrict__ dx, int T, - int H + int H, + float weight_offset ) { int row = blockIdx.x; int tid = threadIdx.x; @@ -168,6 +177,7 @@ __global__ void rmsnorm_bwd_dx_kernel( float dyv = load_as_float(dy_row + col); float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; local_dot += dyv * wv * xv; } @@ -180,6 +190,7 @@ __global__ void rmsnorm_bwd_dx_kernel( float dyv = load_as_float(dy_row + col); float xv = load_as_float(x_row + col); float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; float out = r * dyv * wv - xv * coeff; store_from_float(dx_row + col, out); @@ -187,6 +198,133 @@ __global__ void rmsnorm_bwd_dx_kernel( } +// --------------------------------------------------------------------------- +// Gated RMSNorm (Qwen3-Next GDN block). +// +// y = x * rstd * (weight_offset + weight) * act(gate) +// +// The gate activation and the weight multiply are both evaluated in fp32 and +// there is exactly one cast, at the store. This mirrors vLLM's RMSNormGated +// with norm_before_gate=True, which is how the GDN block constructs it. +// +// ACT selects the gate activation: 0 = silu/swish, 1 = sigmoid. expf (not the +// __expf intrinsic) is used deliberately -- the fast intrinsic trades accuracy +// for speed and would put the result further from the fp32 reference. +// --------------------------------------------------------------------------- + +template +__device__ __forceinline__ float gate_activation(float z) { + const float sigma = 1.0f / (1.0f + expf(-z)); + return (ACT == 0) ? z * sigma : sigma; +} + +template +__device__ __forceinline__ float gate_activation_grad(float z) { + const float sigma = 1.0f / (1.0f + expf(-z)); + // d/dz [z * sigma] = sigma * (1 + z * (1 - sigma)); d/dz [sigma] = sigma * (1 - sigma) + return (ACT == 0) ? sigma * (1.0f + z * (1.0f - sigma)) : sigma * (1.0f - sigma); +} + + +template +__global__ void rmsnorm_gated_fwd_kernel( + const scalar_t* __restrict__ x, + const weight_t* __restrict__ weight, + const scalar_t* __restrict__ gate, + scalar_t* __restrict__ y, + float* __restrict__ rstd, + int T, + int H, + float eps, + float weight_offset +) { + int row = blockIdx.x; + int tid = threadIdx.x; + + const scalar_t* x_row = x + row * H; + const scalar_t* gate_row = gate + row * H; + scalar_t* y_row = y + row * H; + + float local_sum = 0.0f; + + // The statistic is over x only; the gate never enters the reduction, so + // rstd here is bit-identical to the ungated kernel's for the same x. + for (int col = tid; col < H; col += blockDim.x) { + float xv = load_as_float(x_row + col); + local_sum += xv * xv; + } + + float sum = block_reduce_sum(local_sum); + + float row_rstd = rsqrtf(sum / static_cast(H) + eps); + + if (tid == 0) { + rstd[row] = row_rstd; + } + + __syncthreads(); + + for (int col = tid; col < H; col += blockDim.x) { + float xv = load_as_float(x_row + col); + float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; + float zv = load_as_float(gate_row + col); + float out = xv * row_rstd * wv * gate_activation(zv); + store_from_float(y_row + col, out); + } +} + + +template +__global__ void rmsnorm_gated_bwd_dx_kernel( + const scalar_t* __restrict__ dy, + const scalar_t* __restrict__ x, + const weight_t* __restrict__ weight, + const scalar_t* __restrict__ gate, + const float* __restrict__ rstd, + scalar_t* __restrict__ dx, + int T, + int H, + float weight_offset +) { + int row = blockIdx.x; + int tid = threadIdx.x; + + const scalar_t* dy_row = dy + row * H; + const scalar_t* x_row = x + row * H; + const scalar_t* gate_row = gate + row * H; + scalar_t* dx_row = dx + row * H; + + float local_dot = 0.0f; + + // Identical to the ungated dx, with the per-column scale (w + offset) + // replaced by (w + offset) * act(gate): the gate is a constant wrt x. + for (int col = tid; col < H; col += blockDim.x) { + float dyv = load_as_float(dy_row + col); + float xv = load_as_float(x_row + col); + float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; + float zv = load_as_float(gate_row + col); + local_dot += dyv * (wv * gate_activation(zv)) * xv; + } + + float dot = block_reduce_sum(local_dot); + + float r = rstd[row]; + float coeff = dot * r * r * r / static_cast(H); + + for (int col = tid; col < H; col += blockDim.x) { + float dyv = load_as_float(dy_row + col); + float xv = load_as_float(x_row + col); + float wv = load_as_float(weight + col); + if (weight_offset != 0.0f) wv += weight_offset; + float zv = load_as_float(gate_row + col); + + float out = r * dyv * (wv * gate_activation(zv)) - xv * coeff; + store_from_float(dx_row + col, out); + } +} + template __global__ void rmsnorm_partial_dw_kernel( const scalar_t* __restrict__ dy, @@ -247,12 +385,14 @@ void rmsnorm_forward_cuda( torch::Tensor weight, torch::Tensor y, torch::Tensor rstd, - double eps + double eps, + double weight_offset ) { // Launch on x's device: the current CUDA stream belongs to the current // device, which need not be x's. const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); + if (T == 0) return; int H = x.size(1); int threads = choose_threads(H); size_t smem = threads * sizeof(float); @@ -270,10 +410,12 @@ void rmsnorm_forward_cuda( rstd.data_ptr(), T, H, - static_cast(eps) + static_cast(eps), + static_cast(weight_offset) ); }); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } @@ -282,10 +424,12 @@ void rmsnorm_backward_dx_cuda( torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, - torch::Tensor dx + torch::Tensor dx, + double weight_offset ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); + if (T == 0) return; int H = x.size(1); int threads = choose_threads(H); size_t smem = threads * sizeof(float); @@ -303,13 +447,94 @@ void rmsnorm_backward_dx_cuda( rstd.data_ptr(), dx.data_ptr(), T, - H + H, + static_cast(weight_offset) ); }); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + + +void rmsnorm_gated_forward_cuda( + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor y, + torch::Tensor rstd, + double eps, + double weight_offset, + int64_t activation +) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); + int T = x.size(0); + if (T == 0) return; + int H = x.size(1); + int threads = choose_threads(H); + size_t smem = threads * sizeof(float); + + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, x.scalar_type(), "rmsnorm_gated_forward_cuda", [&] { + using x_t = scalar_t; + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, weight.scalar_type(), "rmsnorm_gated_forward_weight_cuda", [&] { + using w_t = scalar_t; + if (activation == 0) { + rmsnorm_gated_fwd_kernel<<>>( + x.data_ptr(), weight.data_ptr(), gate.data_ptr(), + y.data_ptr(), rstd.data_ptr(), + T, H, static_cast(eps), static_cast(weight_offset)); + } else { + rmsnorm_gated_fwd_kernel<<>>( + x.data_ptr(), weight.data_ptr(), gate.data_ptr(), + y.data_ptr(), rstd.data_ptr(), + T, H, static_cast(eps), static_cast(weight_offset)); + } + }); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } +void rmsnorm_gated_backward_dx_cuda( + torch::Tensor dy, + torch::Tensor x, + torch::Tensor weight, + torch::Tensor gate, + torch::Tensor rstd, + torch::Tensor dx, + double weight_offset, + int64_t activation +) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); + int T = x.size(0); + if (T == 0) return; + int H = x.size(1); + int threads = choose_threads(H); + size_t smem = threads * sizeof(float); + + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, x.scalar_type(), "rmsnorm_gated_backward_dx_cuda", [&] { + using x_t = scalar_t; + AT_DISPATCH_FLOATING_TYPES_AND2(at::kHalf, at::kBFloat16, weight.scalar_type(), "rmsnorm_gated_backward_dx_weight_cuda", [&] { + using w_t = scalar_t; + if (activation == 0) { + rmsnorm_gated_bwd_dx_kernel<<>>( + dy.data_ptr(), x.data_ptr(), weight.data_ptr(), + gate.data_ptr(), rstd.data_ptr(), dx.data_ptr(), + T, H, static_cast(weight_offset)); + } else { + rmsnorm_gated_bwd_dx_kernel<<>>( + dy.data_ptr(), x.data_ptr(), weight.data_ptr(), + gate.data_ptr(), rstd.data_ptr(), dx.data_ptr(), + T, H, static_cast(weight_offset)); + } + }); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + void rmsnorm_backward_partial_dw_cuda( torch::Tensor dy, torch::Tensor x, @@ -319,6 +544,7 @@ void rmsnorm_backward_partial_dw_cuda( ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); + if (T == 0) return; int H = x.size(1); int chunks = (T + RMSNORM_DW_ROWS_PER_CHUNK - 1) / RMSNORM_DW_ROWS_PER_CHUNK; @@ -338,6 +564,7 @@ void rmsnorm_backward_partial_dw_cuda( H ); }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } @@ -360,6 +587,7 @@ void rmsnorm_backward_reduce_dw_cuda( chunks, H ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); } #if !defined(USE_ROCM) diff --git a/docs/.nav.yml b/docs/.nav.yml index 29b5ead37..4e04fd061 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -29,6 +29,8 @@ nav: - operators/sampling.md - operators/det-gemm.md - operators/embedding.md + - operators/qwen3-next-rms-norm.md + - operators/qwen3-next-rms-norm-gated.md - Developer Guide: - contributing/README.md - Contributor Guide: contributing/contributor-guide.md diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..0b93b3d90 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -32,4 +32,6 @@ Every operator page should include: - [Matmul](matmul.md) - [Sampling](sampling.md) - [Token Embedding](embedding.md) +- [Qwen3-Next RMSNorm (zero-centred)](qwen3-next-rms-norm.md) +- [Qwen3-Next Gated RMSNorm](qwen3-next-rms-norm-gated.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/qwen3-next-rms-norm-gated.md b/docs/operators/qwen3-next-rms-norm-gated.md new file mode 100644 index 000000000..366f31938 --- /dev/null +++ b/docs/operators/qwen3-next-rms-norm-gated.md @@ -0,0 +1,197 @@ +# 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. + +## Existing implementations: measured batch invariance, accuracy, gates + +![rms_norm_gated vs existing implementations](../usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.png) + +| Implementation | Batch-invariance checks | Forward correctly rounded | Worst grad err | Forward / fwd+bwd | C3/C4 gates | +|---|---|---|---|---|---| +| rl-kernel Qwen3NextRMSNormGatedCudaOp | yes (sampled) | 99.9989% | 2.4e-03 | 235 / 115293 µs | pass | +| rl-kernel PyTorch reference | yes (sampled) | 99.9990% | 2.4e-03 | 772 / 1991 µs | — | +| transformers 5.17.0 Qwen3NextRMSNormGated | yes (sampled) | 65.4748% | 5.6e-03 | 695 / 2093 µs | pass | +| FLA 0.5.2 layernorm_gated.rmsnorm_fn | yes (sampled) | 99.9987% | 2.4e-03 | 177 / 882 µs | pass | +| FLA 0.5.2 fused_norm_gate.rms_norm_gated | **no** (247 sub-batches) | 99.9987% | 2.4e-03 | 93 / 865 µs | pass | +| vLLM 0.30.0 RMSNormGated.forward_cuda | **no** (3 rows, 3 sub-batches) | 99.9988% | — | 87 / — µs | — | + +Batch invariance is bitwise and covers three checks: every row computed alone vs inside full +batches of three sizes; sampled small sub-batches and exhaustive larger sub-batches +of the full workload-size batch; and a batch-size sweep over probe rows. In these +historical reports, sub-batch sizes 1 and 7 visit only 512 starts per seed, so a +"yes (sampled)" does not establish every-row coverage at those sizes. The JSON +records the actual stride and rows checked per size. The current script instead +partitions the full batch at every advertised sub-batch size, covering every row +including a final partial batch. These historical measurements have not been rerun. +A "no" counts the rows, sub-batches or sweep cases that differed. Accuracy is against FP64 at the workload size; latency is the median on an otherwise +idle B200. The C3/C4 column runs this repository's own gate scripts unchanged, with the CUDA +candidate replaced by a subclass of this op whose forward and backward call the other library. +The subclass keeps this op's FP32 `dweight` row contributions, so singleton-aggregate compares +like with like. [`rms_norm_gated.json`](../usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.json) +was written from a clean tree at `abb56c2` by + +```bash +python tools/validation/models/qwen3_next_norm_reuse_check.py --op rms_norm_gated \ + --out docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.json [--megatron-src ] +``` + +## 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..945c8ef7a --- /dev/null +++ b/docs/operators/qwen3-next-rms-norm.md @@ -0,0 +1,198 @@ +# 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 wrapper copies non-contiguous inputs; low-level `rmsnorm_cuda` requires contiguous inputs | +| `weight` | `[H]` | matches `x` | zero-centred (upstream inits to zeros); CUDA wrapper copies non-contiguous weights; low-level `rmsnorm_cuda` requires contiguous weights | +| `eps` | scalar | float | inside the sqrt; `1e-6` for Qwen3-Next | + +## Dispatch Behavior + +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. + +## Existing implementations: measured batch invariance, accuracy, gates + +![qwen3_next_rms_norm vs existing implementations](../usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png) + +| Implementation | Batch-invariance checks | Forward correctly rounded | Worst grad err | Forward / fwd+bwd | C3/C4 gates | +|---|---|---|---|---|---| +| rl-kernel Qwen3NextRMSNormCudaOp | yes (sampled) | 99.9993% | 2.4e-03 | 426 / 31485 µs | see gate table below | +| rl-kernel PyTorch reference | **no** (49 rows, 94 sub-batches) | 99.9993% | 2.4e-03 | 2146 / 5024 µs | — | +| transformers 5.17.0 Qwen3NextRMSNorm | **no** (90 rows, 142 sub-batches) | 99.9992% | 2.4e-03 | 1449 / 4468 µs | see gate table below | +| Liger 0.8.4 RMSNorm, offset 1, gemma | yes (sampled) | 99.9993% | 2.4e-03 | 133 / 974 µs | see gate table below | +| FLA 0.5.2 rms_norm, weight passed as 1 + w | yes (sampled) | 73.0161% | 5.5e-03 | 147 / 1051 µs | see gate table below | +| TE 2.20.2 RMSNorm(zero_centered_gamma) | **no** (175 rows, 245 sub-batches) | 99.9992% | 2.4e-03 | 156 / 779 µs | — | +| FlashInfer 0.6.18.post1 gemma_rmsnorm | yes (sampled) | 99.9992% | — | 91 / — µs | — | +| vLLM 0.30.0 GemmaRMSNorm.forward_cuda | **no** (16 rows, 26 sub-batches) | 99.9992% | — | 1454 / — µs | — | + +Batch invariance is bitwise and covers three checks: every row computed alone vs inside full +batches of three sizes; sampled small sub-batches and exhaustive larger sub-batches +of the full workload-size batch; and a batch-size sweep over probe rows. In these +historical reports, sub-batch sizes 1 and 7 visit only 512 starts per seed, so a +"yes (sampled)" does not establish every-row coverage at those sizes. The JSON +records the actual stride and rows checked per size. The current script instead +partitions the full batch at every advertised sub-batch size, covering every row +including a final partial batch. These historical measurements have not been rerun. +A "no" counts the rows, sub-batches or sweep cases that differed. Accuracy is against FP64 at the workload size; latency is the median on an otherwise +idle B200. The C3/C4 column runs this repository's own gate scripts unchanged, with the CUDA +candidate replaced by a subclass of this op whose forward and backward call the other library. +The subclass keeps this op's FP32 `dweight` row contributions, so singleton-aggregate compares +like with like. The C3/C4 results for this op are in the gate table below. [`qwen3_next_rms_norm.json`](../usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json) +was originally written from a clean tree at `a66493c` by the command below. The +Megatron `52fbcbc` result has been excluded from the reports, tables and figure: +that revision computes `weight_eff` but uses the original `weight` in its forward +output, so it does not implement the zero-centred operation despite accepting the +flag. The remaining measurements are unchanged; the figure was regenerated from +the corrected report without rerunning benchmarks. The reuse checker now rejects +Megatron implementations that fail a zero-weight probe before benchmarking them. + +Original command: + +```bash +python tools/validation/models/qwen3_next_norm_reuse_check.py --op qwen3_next_rms_norm \ + --out docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json [--megatron-src ] +``` + + +The gate scripts support this op on this branch. Run unchanged at `88f59f7`, with the CUDA candidate swapped for each implementation ([`qwen3_next_rms_norm_gates.json`](../usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm_gates.json), `tools/validation/models/qwen3_next_norm_reuse_check.py --op qwen3_next_rms_norm --checks gates`): + +| Implementation swapped into the CUDA candidate | C3 forward | C4 gradient (incl. singleton-aggregate) | +|---|---|---| +| rl-kernel Qwen3NextRMSNormCudaOp | pass | pass | +| transformers 5.17.0 Qwen3NextRMSNorm | pass | **fail** | +| Liger 0.8.4 RMSNorm, offset 1, gemma | pass | pass | +| FLA 0.5.2 rms_norm, weight passed as 1 + w | pass | pass | + +## 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-norm-reuse-b200/qwen3_next_rms_norm.json b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json new file mode 100644 index 000000000..9b3b2f8b6 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json @@ -0,0 +1,1043 @@ +{ + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": "qwen3_next_rms_norm", + "rl_kernel_commit": "a66493c6d0b49d78aec6df3519b70fb45eace24f", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "cuda": "13.0", + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "liger-kernel": "0.8.4", + "fla-core": "0.5.2", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc" + }, + "quick": false, + "checks": [ + "bi", + "gates", + "perf" + ], + "results": { + "rl_kernel": { + "source": "rl-kernel Qwen3NextRMSNormCudaOp", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999926686286926, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 426.4480024576187, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 31485.487937927246 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "reference": { + "source": "rl-kernel PyTorch reference", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 2 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 7 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 40 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 94, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999927282333374, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 2145.951986312866, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 5024.384021759033 + } + }, + "transformers": { + "source": "transformers 5.17.0 Qwen3NextRMSNorm", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 3 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 10 + }, + "2048": { + "rows": 6144, + "fwd_differ": 16, + "rowgrad_differ": 61 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 26, + "rowgrad_differ": 116, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999924898147583, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 1449.1679668426514, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 4467.727899551392 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "liger": { + "source": "Liger 0.8.4 RMSNorm, offset 1, gemma", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999925494194031, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 133.31200182437897, + "grad_rel_err": { + "x": 0.0023939748586017722, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 973.6000001430511 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "fla": { + "source": "FLA 0.5.2 rms_norm, weight passed as 1 + w", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.730160653591156, + "fwd_max_abs_err": 0.03517675224226302, + "fwd_us": 146.92799746990204, + "grad_rel_err": { + "x": 0.005450811301370503, + "w": 0.0019861680402388695 + }, + "fwd_bwd_us": 1050.8000254631042 + }, + "contract_gates": { + "unavailable": "rl_engine/testing/qwen3_next_norm_manifest.json is not on this branch" + } + }, + "transformer_engine": { + "source": "TE 2.20.2 RMSNorm(zero_centered_gamma)", + "backward": "x-only", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 1, + "rowgrad_differ": 2 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 75, + "rowgrad_differ": 97 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 111, + "rowgrad_differ": 134, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999922513961792, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 156.40000253915787, + "grad_rel_err": { + "x": 0.0023939340856740927 + }, + "fwd_bwd_us": 778.5120010375977 + } + }, + "megatron": { + "unavailable": "Excluded: Megatron 52fbcbc accepts zero_centered_gamma=True but uses the original weight instead of weight_eff in forward." + }, + "flashinfer": { + "source": "FlashInfer 0.6.18.post1 gemma_rmsnorm", + "backward": "none", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999917149543762, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 90.71999788284302 + } + }, + "vllm": { + "source": "vLLM 0.30.0 GemmaRMSNorm.forward_cuda", + "backward": "none", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 16, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 65536, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097, + 32768 + ], + "sub_batches": 2148, + "fwd_differ": 26, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 128, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 128, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "2048": { + "stride": 2048, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 65536, + "covers_every_row": true + }, + "32768": { + "stride": 32768, + "rows_checked_per_seed": 65536, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2232, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 65536, + "fwd_correctly_rounded": 0.9999924898147583, + "fwd_max_abs_err": 0.015624080561254416, + "fwd_us": 1454.3840289115906 + } + } + }, + "evidence_corrections": [ + "Removed the invalid Megatron 52fbcbc zero-centered comparison and regenerated the figure. All other measurements and original run provenance are unchanged; benchmarks were not rerun.", + "Annotated historical sub-batch sampling coverage and regenerated the figure labels. Measurements were not rerun; batch_invariant records only whether the measured checks passed, not exhaustive full-batch coverage at every sub-batch size." + ] +} diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png new file mode 100644 index 000000000..3f7a645aa Binary files /dev/null and b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.png differ diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm_gates.json b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm_gates.json new file mode 100644 index 000000000..7d3aee03e --- /dev/null +++ b/docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm_gates.json @@ -0,0 +1,152 @@ +{ + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": "qwen3_next_rms_norm", + "rl_kernel_commit": "88f59f7d06741ba533a7b9b718b331ee885b93b7", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "cuda": "13.0", + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "liger-kernel": "0.8.4", + "fla-core": "0.5.2", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc" + }, + "quick": false, + "checks": [ + "gates" + ], + "results": { + "rl_kernel": { + "source": "rl-kernel Qwen3NextRMSNormCudaOp", + "backward": "full", + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "reference": { + "source": "rl-kernel PyTorch reference", + "backward": "full" + }, + "transformers": { + "source": "transformers 5.17.0 Qwen3NextRMSNorm", + "backward": "full", + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 1, + "passed": false, + "failed_lines": [ + "op=qwen3_next_rms_norm profile=cuda_bf16 candidate=__main__._gate_child..Swapped::cuda passed=False", + "invariance pair=('BN/full', 'BN/chunked') tensor=dx transform=chunk max_abs=1.52587891e-05 passed=False", + "invariance pair=('BN/full', 'B1-singleton_aggregate/full/s2') tensor=dx transform=batch_size max_abs=1.52587891e-05 passed=False", + "invariance pair=('BN/full', 'B1-singleton_aggregate/chunked/s2') tensor=dx transform=chunk max_abs=1.52587891e-05 passed=False" + ], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "liger": { + "source": "Liger 0.8.4 RMSNorm, offset 1, gemma", + "backward": "full", + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "fla": { + "source": "FLA 0.5.2 rms_norm, weight passed as 1 + w", + "backward": "full", + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "transformer_engine": { + "source": "TE 2.20.2 RMSNorm(zero_centered_gamma)", + "backward": "x-only" + }, + "megatron": { + "unavailable": "Excluded: Megatron 52fbcbc accepts zero_centered_gamma=True but uses the original weight instead of weight_eff in forward." + }, + "flashinfer": { + "source": "FlashInfer 0.6.18.post1 gemma_rmsnorm", + "backward": "none" + }, + "vllm": { + "source": "vLLM 0.30.0 GemmaRMSNorm.forward_cuda", + "backward": "none" + } + }, + "evidence_corrections": [ + "Removed the invalid Megatron 52fbcbc zero-centered gate comparison. All other gate measurements and original run provenance are unchanged; gates were not rerun." + ] +} diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.json b/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.json new file mode 100644 index 000000000..397879351 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.json @@ -0,0 +1,888 @@ +{ + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": "rms_norm_gated", + "rl_kernel_commit": "abb56c2d17d36d44efbe2e7ecffe87602496083b", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "cuda": "13.0", + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "liger-kernel": "0.8.4", + "fla-core": "0.5.2", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc" + }, + "quick": false, + "checks": [ + "bi", + "gates", + "perf" + ], + "results": { + "rl_kernel": { + "source": "rl-kernel Qwen3NextRMSNormGatedCudaOp", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.9999886155128479, + "fwd_max_abs_err": 0.031235785491551482, + "fwd_us": 234.6400022506714, + "grad_rel_err": { + "x": 0.0020048456424744273, + "z": 0.002181870606557371, + "w": 0.002359713325541337 + }, + "fwd_bwd_us": 115293.25103759766 + }, + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "reference": { + "source": "rl-kernel PyTorch reference", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.9999895691871643, + "fwd_max_abs_err": 0.031235785491551482, + "fwd_us": 771.6160118579865, + "grad_rel_err": { + "x": 0.0020048456424744273, + "z": 0.002181870606557371, + "w": 0.002359713325541337 + }, + "fwd_bwd_us": 1991.0719990730286 + } + }, + "transformers": { + "source": "transformers 5.17.0 Qwen3NextRMSNormGated", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.6547477841377258, + "fwd_max_abs_err": 0.08551652812757382, + "fwd_us": 694.7199702262878, + "grad_rel_err": { + "x": 0.005555031181417614, + "z": 0.0051996228630245295, + "w": 0.004407626951010427 + }, + "fwd_bwd_us": 2092.736005783081 + }, + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "fla_layernorm_gated": { + "source": "FLA 0.5.2 layernorm_gated.rmsnorm_fn", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 0, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": true + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.9999874830245972, + "fwd_max_abs_err": 0.031235785491551482, + "fwd_us": 176.70400440692902, + "grad_rel_err": { + "x": 0.0020048456424744273, + "z": 0.002181870606557371, + "w": 0.002359713325541337 + }, + "fwd_bwd_us": 882.4159801006317 + }, + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "fla_fused_norm_gate": { + "source": "FLA 0.5.2 fused_norm_gate.rms_norm_gated", + "backward": "full", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 95, + "rowgrad_differ": 152, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.9999873638153076, + "fwd_max_abs_err": 0.031235785491551482, + "fwd_us": 93.37600320577621, + "grad_rel_err": { + "x": 0.0020048456424744273, + "z": 0.002181870606557371, + "w": 0.002359713325541337 + }, + "fwd_bwd_us": 865.2800023555756 + }, + "contract_gates": { + "forward": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [], + "stderr_tail": [] + }, + "gradient": { + "returncode": 0, + "passed": true, + "failed_lines": [], + "singleton_aggregate": [ + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/full') tensor=dweight max_abs=0.00000000e+00 passed=True", + "singleton_aggregate pair=('BN/full', 'B1-singleton_aggregate/chunked') tensor=dweight max_abs=0.00000000e+00 passed=True" + ], + "stderr_tail": [] + } + } + }, + "vllm": { + "source": "vLLM 0.30.0 RMSNormGated.forward_cuda", + "backward": "none", + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "4097": { + "rows": 12291, + "fwd_differ": 3, + "rowgrad_differ": 0 + } + }, + "dweight_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 262144, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 65536, + 131072 + ], + "sub_batches": 2188, + "fwd_differ": 3, + "rowgrad_differ": 0, + "coverage": "sampled_small_sub_batches", + "coverage_by_sub_batch_size": { + "1": { + "stride": 512, + "rows_checked_per_seed": 512, + "covers_every_row": false + }, + "7": { + "stride": 512, + "rows_checked_per_seed": 3584, + "covers_every_row": false + }, + "4097": { + "stride": 4097, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "65536": { + "stride": 65536, + "rows_checked_per_seed": 262144, + "covers_every_row": true + }, + "131072": { + "stride": 131072, + "rows_checked_per_seed": 262144, + "covers_every_row": true + } + } + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048, + 2049, + 4095, + 4096, + 4097 + ], + "probe_rows": [ + 0, + 515, + 1027, + 1539, + 2051, + 2563, + 3075, + 4096 + ], + "comparisons": 2520, + "differ": 0 + }, + "batch_invariant": false + }, + "accuracy_latency": { + "rows": 262144, + "fwd_correctly_rounded": 0.9999884366989136, + "fwd_max_abs_err": 0.031235785491551482, + "fwd_us": 87.07199990749359 + } + } + }, + "evidence_corrections": [ + "Annotated historical sub-batch sampling coverage and regenerated the figure labels. Measurements were not rerun; batch_invariant records only whether the measured checks passed, not exhaustive full-batch coverage at every sub-batch size." + ] +} diff --git a/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.png b/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.png new file mode 100644 index 000000000..a6ed8310c Binary files /dev/null and b/docs/usage/evidence/qwen3-next-norm-reuse-b200/rms_norm_gated.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png b/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png new file mode 100644 index 000000000..a408aff9c Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-b200/figure.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json b/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json new file mode 100644 index 000000000..0fc8452ae --- /dev/null +++ b/docs/usage/evidence/qwen3-next-rms-norm-b200/report.json @@ -0,0 +1,225 @@ +{ + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": "3d0bae7174c09a133ea1288d820b63e4a7db8f57", + "git_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer": "0.6.18.post1" + }, + "hidden": 2048, + "eps": 1e-06, + "dtype": "bfloat16", + "ops": { + "zero_centred_rmsnorm": { + "accuracy": { + "257": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999886155128479 + } + }, + "4096": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999991774559021 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999927282333374 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999992847442627 + } + } + }, + "row_invariance": { + "rl-kernel CUDA": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "rl-kernel PyTorch reference": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 12, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "transformers Qwen3NextRMSNorm": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": 11, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "vLLM GemmaRMSNorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + } + }, + "latency": { + "rl-kernel CUDA": { + "1024": { + "forward_us": 24.70399998128414, + "backward_us": 568.7519907951355 + }, + "4096": { + "forward_us": 39.16800022125244, + "backward_us": 1259.3119740486145 + }, + "16384": { + "forward_us": 122.36800044775009, + "backward_us": 7879.024028778076 + }, + "65536": { + "forward_us": 425.80799758434296, + "backward_us": 30937.551498413086 + } + }, + "rl-kernel PyTorch reference": { + "1024": { + "forward_us": 70.54400071501732, + "backward_us": 644.7039842605591 + }, + "4096": { + "forward_us": 155.90400248765945, + "backward_us": 624.2719888687134 + }, + "16384": { + "forward_us": 577.5039792060852, + "backward_us": 986.7520034313202 + }, + "65536": { + "forward_us": 2148.0319499969482, + "backward_us": 2990.224003791809 + } + }, + "transformers Qwen3NextRMSNorm": { + "1024": { + "forward_us": 69.34399902820587, + "backward_us": 608.5599958896637 + }, + "4096": { + "forward_us": 117.3119992017746, + "backward_us": 606.3359975814819 + }, + "16384": { + "forward_us": 407.50400722026825, + "backward_us": 1027.888000011444 + }, + "65536": { + "forward_us": 1449.4240283966064, + "backward_us": 3144.6080207824707 + } + }, + "vLLM GemmaRMSNorm (forward only)": { + "1024": { + "forward_us": 65.61600044369698 + }, + "4096": { + "forward_us": 118.40000003576279 + }, + "16384": { + "forward_us": 405.90400993824005 + }, + "65536": { + "forward_us": 1449.6000409126282 + } + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "1024": { + "forward_us": 14.448000118136406 + }, + "4096": { + "forward_us": 16.24000072479248 + }, + "16384": { + "forward_us": 32.207999378442764 + }, + "65536": { + "forward_us": 90.55999666452408 + } + } + } + } + } +} diff --git a/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png new file mode 100644 index 000000000..eaf59c418 Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-gated_rmsnorm.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png new file mode 100644 index 000000000..5cd70b5cf Binary files /dev/null and b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/figure-zero_centred_rmsnorm.png differ diff --git a/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json new file mode 100644 index 000000000..8a931b315 --- /dev/null +++ b/docs/usage/evidence/qwen3-next-rms-norm-gated-b200/report.json @@ -0,0 +1,400 @@ +{ + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": "822b085cd5b042661b2bda021414a2caa796bef0", + "git_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "transformers": "5.17.0", + "vllm": "0.30.0", + "flashinfer": "0.6.18.post1" + }, + "eps": 1e-06, + "dtype": "bfloat16", + "ops": { + "zero_centred_rmsnorm": { + "hidden": 2048, + "accuracy": { + "257": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.002701418159021442, + "dweight_max_abs_over_absmax": 0.0023356239188215265, + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999905228614807 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.01552825644384992, + "forward_correctly_rounded_fraction": 0.9999886155128479 + } + }, + "4096": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999991774559021 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999927282333374 + }, + "transformers Qwen3NextRMSNorm": { + "dx_max_abs_over_absmax": 0.0027934846392837836, + "dweight_max_abs_over_absmax": 0.002058919545322908, + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "vLLM GemmaRMSNorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.9999921321868896 + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "forward_max_abs": 0.015599696412158082, + "forward_correctly_rounded_fraction": 0.999992847442627 + } + } + }, + "row_invariance": { + "rl-kernel CUDA": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "rl-kernel PyTorch reference": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 12, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "transformers Qwen3NextRMSNorm": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": 11, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "vLLM GemmaRMSNorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 2, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + } + }, + "latency": { + "rl-kernel CUDA": { + "1024": { + "forward_us": 25.200000032782555, + "backward_us": 556.6880106925964 + }, + "4096": { + "forward_us": 39.503999054431915, + "backward_us": 1260.4320049285889 + }, + "16384": { + "forward_us": 122.41600081324577, + "backward_us": 7879.663944244385 + }, + "65536": { + "forward_us": 427.5680035352707, + "backward_us": 30925.631523132324 + } + }, + "rl-kernel PyTorch reference": { + "1024": { + "forward_us": 71.45600020885468, + "backward_us": 595.3599810600281 + }, + "4096": { + "forward_us": 156.92799538373947, + "backward_us": 537.9360020160675 + }, + "16384": { + "forward_us": 578.4800052642822, + "backward_us": 992.5920069217682 + }, + "65536": { + "forward_us": 2150.12788772583, + "backward_us": 2983.504056930542 + } + }, + "transformers Qwen3NextRMSNorm": { + "1024": { + "forward_us": 70.25599852204323, + "backward_us": 599.3280112743378 + }, + "4096": { + "forward_us": 118.12799796462059, + "backward_us": 541.1999821662903 + }, + "16384": { + "forward_us": 409.2479944229126, + "backward_us": 1025.5680084228516 + }, + "65536": { + "forward_us": 1452.351987361908, + "backward_us": 3140.112042427063 + } + }, + "vLLM GemmaRMSNorm (forward only)": { + "1024": { + "forward_us": 67.4239993095398 + }, + "4096": { + "forward_us": 119.00799721479416 + }, + "16384": { + "forward_us": 407.3439985513687 + }, + "65536": { + "forward_us": 1451.9200325012207 + } + }, + "FlashInfer gemma_rmsnorm (forward only)": { + "1024": { + "forward_us": 14.944000169634819 + }, + "4096": { + "forward_us": 17.008000053465366 + }, + "16384": { + "forward_us": 32.94399939477444 + }, + "65536": { + "forward_us": 90.43200314044952 + } + } + } + }, + "gated_rmsnorm": { + "hidden": 128, + "accuracy": { + "257": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0018348192997423702, + "dweight_max_abs_over_absmax": 0.0021439847223899285, + "dgate_max_abs_over_absmax": 0.0027467772330278155, + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0018348192997423702, + "dweight_max_abs_over_absmax": 0.0021439847223899285, + "dgate_max_abs_over_absmax": 0.0027467772330278155, + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "dx_max_abs_over_absmax": 0.004283006956510838, + "dweight_max_abs_over_absmax": 0.0053612828842400685, + "dgate_max_abs_over_absmax": 0.004464247864523204, + "forward_max_abs": 0.041932799726739134, + "forward_correctly_rounded_fraction": 0.6546084880828857 + }, + "vLLM RMSNormGated (forward only)": { + "forward_max_abs": 0.01814649124332668, + "forward_correctly_rounded_fraction": 0.9999392032623291 + } + }, + "4096": { + "rl-kernel CUDA": { + "dx_max_abs_over_absmax": 0.0028905074752616756, + "dweight_max_abs_over_absmax": 0.0017551573197679504, + "dgate_max_abs_over_absmax": 0.002440494893713147, + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "rl-kernel PyTorch reference": { + "dx_max_abs_over_absmax": 0.0028905074752616756, + "dweight_max_abs_over_absmax": 0.0017551573197679504, + "dgate_max_abs_over_absmax": 0.002440494893713147, + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999942779541016 + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "dx_max_abs_over_absmax": 0.006106839188157589, + "dweight_max_abs_over_absmax": 0.002835271706901685, + "dgate_max_abs_over_absmax": 0.003959744684751284, + "forward_max_abs": 0.06659807537061546, + "forward_correctly_rounded_fraction": 0.6562480926513672 + }, + "vLLM RMSNormGated (forward only)": { + "forward_max_abs": 0.028763107099734952, + "forward_correctly_rounded_fraction": 0.9999923706054688 + } + } + }, + "row_invariance": { + "rl-kernel CUDA": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "rl-kernel PyTorch reference": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": 0, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + }, + "vLLM RMSNormGated (forward only)": { + "rows_checked": 768, + "forward_rows_differing": 0, + "dx_rows_differing": null, + "batch_rows": 4096, + "seeds": [ + 3, + 4, + 5 + ] + } + }, + "latency": { + "rl-kernel CUDA": { + "4096": { + "forward_us": 20.800000056624413, + "backward_us": 1269.7439789772034 + }, + "16384": { + "forward_us": 27.168000116944313, + "backward_us": 3861.135959625244 + }, + "65536": { + "forward_us": 64.00000303983688, + "backward_us": 14388.879776000977 + }, + "262144": { + "forward_us": 235.23200303316116, + "backward_us": 114274.65438842773 + } + }, + "rl-kernel PyTorch reference": { + "4096": { + "forward_us": 81.31200075149536, + "backward_us": 560.6080293655396 + }, + "16384": { + "forward_us": 89.72799777984619, + "backward_us": 550.2559840679169 + }, + "65536": { + "forward_us": 209.3760073184967, + "backward_us": 573.7600028514862 + }, + "262144": { + "forward_us": 774.8000025749207, + "backward_us": 1376.08003616333 + } + }, + "transformers Qwen3NextRMSNormGated (cast-first)": { + "4096": { + "forward_us": 90.01599997282028, + "backward_us": 577.888011932373 + }, + "16384": { + "forward_us": 91.80799871683121, + "backward_us": 601.1359989643097 + }, + "65536": { + "forward_us": 196.46400213241577, + "backward_us": 623.5039830207825 + }, + "262144": { + "forward_us": 697.2479820251465, + "backward_us": 1528.384029865265 + } + }, + "vLLM RMSNormGated (forward only)": { + "4096": { + "forward_us": 49.775999039411545 + }, + "16384": { + "forward_us": 51.024001091718674 + }, + "65536": { + "forward_us": 50.20799860358238 + }, + "262144": { + "forward_us": 81.4880020916462 + } + } + } + } + } +} diff --git a/rl_engine/_C.pyi b/rl_engine/_C.pyi index fb3c4379c..c0f9cb02c 100644 --- a/rl_engine/_C.pyi +++ b/rl_engine/_C.pyi @@ -256,16 +256,38 @@ def swiglu_backward( gate: torch.Tensor, up: torch.Tensor, ) -> list[torch.Tensor]: ... + +rmsnorm_api_version: int + def rmsnorm_forward( x: torch.Tensor, weight: torch.Tensor, eps: float, + weight_offset: float = ..., ) -> list[torch.Tensor]: ... def rmsnorm_backward_dx( dy: torch.Tensor, x: torch.Tensor, weight: torch.Tensor, rstd: torch.Tensor, + weight_offset: float = ..., +) -> torch.Tensor: ... +def rmsnorm_gated_forward( + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + eps: float, + weight_offset: float = ..., + activation: int = ..., +) -> list[torch.Tensor]: ... +def rmsnorm_gated_backward_dx( + dy: torch.Tensor, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + rstd: torch.Tensor, + weight_offset: float = ..., + activation: int = ..., ) -> torch.Tensor: ... def rmsnorm_backward_dw( dy: torch.Tensor, diff --git a/rl_engine/backends/cuda/norm/rmsnorm.py b/rl_engine/backends/cuda/norm/rmsnorm.py index ec3c65f77..60cdbffd3 100644 --- a/rl_engine/backends/cuda/norm/rmsnorm.py +++ b/rl_engine/backends/cuda/norm/rmsnorm.py @@ -4,9 +4,21 @@ from rl_engine.ops.autograd.backward_runtime import record_backward from rl_engine.ops.autograd.vjp_fp32 import reduce_rows_fp32, rmsnorm_dweight_rows_fp32 +_RMSNORM_API_VERSION = 2 -def _require_cuda_symbols(what: str, *names: str) -> None: - """Raise when the compiled kernels backing ``what`` are missing. + +def _fold_dweight_rows(rows: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: + """The single left-fold entrypoint for this backend's dweight reductions. + + Both the plain and the gated backward route through here so the file keeps + one auditable reduction path; the ascending-row fp32 fold is what makes + dweight independent of the batch layout. + """ + return reduce_rows_fp32(rows).to(dtype) + + +def _require_cuda_rmsnorm() -> None: + """Raise when the compiled RMSNorm bindings are missing or incompatible. The registry treats a backend whose construction raises as unavailable and falls through to the next candidate, so calling this from ``__init__`` is @@ -15,13 +27,21 @@ def _require_cuda_symbols(what: str, *names: str) -> None: activation ops. """ if not _EXT_AVAILABLE or _C is None: - raise RuntimeError(f"{what} requires the compiled rl_engine._C extension.") + raise RuntimeError("CUDA RMSNorm requires the compiled rl_engine._C extension.") + names = ("rmsnorm_forward", "rmsnorm_backward_dx") missing = [name for name in names if not hasattr(_C, name)] if missing: raise RuntimeError( - f"{what} symbols ({', '.join(missing)}) are not compiled into _C. " + f"CUDA RMSNorm symbols ({', '.join(missing)}) are not compiled into _C. " "Rebuild the extension with csrc/cuda/norm/rmsnorm.cu." ) + api_version = getattr(_C, "rmsnorm_api_version", None) + if api_version != _RMSNORM_API_VERSION: + raise RuntimeError( + f"CUDA RMSNorm requires rl_engine._C RMSNorm API version {_RMSNORM_API_VERSION} " + f"(loaded {api_version!r}). Rebuild the extension with csrc/cuda/norm/rmsnorm.cu " + "for weight_offset support." + ) class RMSNormCuda(torch.autograd.Function): @@ -30,20 +50,23 @@ class RMSNormCuda(torch.autograd.Function): """ @staticmethod - def forward(ctx, x, weight, mask=None, eps=1e-6): + def forward(ctx, x, weight, mask=None, eps=1e-6, weight_offset=0.0): """ Forward: - y = x * rsqrt(mean(x^2) + eps) * weight + y = x * rsqrt(mean(x^2) + eps) * (weight_offset + weight) Input: x: [T, H], fp16/bf16/fp32 CUDA tensor weight: [H], fp16/bf16/fp32 CUDA tensor mask: [T], bool CUDA tensor eps: float + weight_offset: float, added to weight in fp32 inside the kernel. + 1.0 selects the zero-centred (1 + w) convention. Output: y: [T, H] """ + _require_cuda_rmsnorm() assert x.is_cuda, "x must be CUDA tensor" assert weight.is_cuda, "weight must be CUDA tensor" assert x.is_contiguous(), "x must be contiguous" @@ -51,10 +74,6 @@ def forward(ctx, x, weight, mask=None, eps=1e-6): assert x.dim() == 2, "x must be [T, H]" assert weight.dim() == 1, "weight must be [H]" assert x.shape[1] == weight.shape[0], "hidden size mismatch" - assert _EXT_AVAILABLE and hasattr( - _C, "rmsnorm_forward" - ), "RMSNorm CUDA extension is unavailable. Please rebuild with rmsnorm.cu." - if mask is None: mask = torch.ones((x.shape[0],), device=x.device, dtype=torch.bool) else: @@ -64,10 +83,11 @@ def forward(ctx, x, weight, mask=None, eps=1e-6): assert mask.dim() == 1, "mask must be [T]" assert mask.shape[0] == x.shape[0], "mask length mismatch" - y, rstd = _C.rmsnorm_forward(x, weight, float(eps)) + y, rstd = _C.rmsnorm_forward(x, weight, float(eps), float(weight_offset)) ctx.save_for_backward(x, weight, rstd, mask) ctx.eps = eps + ctx.weight_offset = float(weight_offset) return y @@ -81,7 +101,7 @@ def backward(ctx, grad_out): x, weight, rstd, mask = ctx.saved_tensors dy = grad_out.contiguous() - dx = _C.rmsnorm_backward_dx(dy, x, weight, rstd) + dx = _C.rmsnorm_backward_dx(dy, x, weight, rstd, ctx.weight_offset) # The shape-independent FP32 left fold preserves the C2 Batch/Chunk # reduction order while the CUDA reducer executes it in one launch. @@ -89,7 +109,7 @@ def backward(ctx, grad_out): # Multiplication is part of the pre-existing mask contract, including # IEEE propagation for non-finite inactive contributions. rows = rows * mask.to(dtype=rows.dtype).unsqueeze(-1) - dw = reduce_rows_fp32(rows).to(weight.dtype) + dw = _fold_dweight_rows(rows, weight.dtype) record_backward( "rms_norm", kernel_id=( @@ -101,16 +121,18 @@ def backward(ctx, grad_out): family="cuda", ) - return dx, dw, None, None + # dw is unchanged by the offset: d/dw (offset + w) == d/dw w. + return dx, dw, None, None, None -def rmsnorm_cuda(x, weight, eps=1e-6, mask=None): +def rmsnorm_cuda(x, weight, eps=1e-6, mask=None, weight_offset=0.0): """ use: y = rmsnorm_cuda(x, weight) y = rmsnorm_cuda(x, weight, mask=mask) + y = rmsnorm_cuda(x, weight, weight_offset=1.0) # zero-centred weight """ - return RMSNormCuda.apply(x, weight, mask, eps) + return RMSNormCuda.apply(x, weight, mask, eps, weight_offset) class RMSNormCudaOp: @@ -119,11 +141,11 @@ class RMSNormCudaOp: backward_impl = "cuda_rmsnorm_dx_declared_fp32_rowfold_dw" def __init__(self) -> None: - _require_cuda_symbols( - "CUDA RMSNorm", - "rmsnorm_forward", - "rmsnorm_backward_dx", - ) + _require_cuda_rmsnorm() + + #: Added to the weight in fp32 inside the kernel. Subclasses override it; + #: 0.0 is the plain convention. + weight_offset = 0.0 def __call__(self, x, weight, *, eps=1e-6): return self.forward(x, weight, eps=eps) @@ -131,12 +153,236 @@ def __call__(self, x, weight, *, eps=1e-6): def forward(self, x, weight, *, eps=1e-6): hidden = x.shape[-1] x_2d = x.contiguous().view(-1, hidden) - y_2d = rmsnorm_cuda(x_2d, weight.contiguous(), eps=eps) + y_2d = rmsnorm_cuda(x_2d, weight.contiguous(), eps=eps, weight_offset=self.weight_offset) return y_2d.view_as(x) def parameter_vjp_contributions_fp32(self, *, x, weight, grad_output, eps=1e-6): - del weight - x32 = x.float() - rstd = torch.rsqrt(x32.square().mean(dim=-1) + float(eps)) - rows = grad_output.float() * x32 * rstd.unsqueeze(-1) + hidden = x.shape[-1] + # Only `rstd` is used, and it does not depend on the offset; the offset is + # passed so this is the same call the forward makes, not because it matters. + _, rstd = _C.rmsnorm_forward( + x.contiguous().reshape(-1, hidden), + weight.contiguous(), + float(eps), + float(self.weight_offset), + ) + rows = rmsnorm_dweight_rows_fp32(x, grad_output, rstd=rstd.reshape(x.shape[:-1])) + return {"weight": rows} + + +class Qwen3NextRMSNormCudaOp(RMSNormCudaOp): + """Zero-centred CUDA RMSNorm: ``y = x * rstd * (1 + weight)``. + + The decoder and final norms of Qwen3-Next (and Gemma) store a zero-centred + weight. The ``+1`` is applied inside the kernel after the fp32 upcast, so it + is never rounded through the low-precision weight dtype. + """ + + weight_offset = 1.0 + + +# --------------------------------------------------------------------------- # +# Gated RMSNorm (Qwen3-Next GDN block) +# --------------------------------------------------------------------------- # + + +def _require_cuda_symbols(what: str, *names: str) -> None: + """Raise when the compiled kernels backing ``what`` are missing. + + The registry treats a backend whose construction raises as unavailable and + falls through, so calling this from ``__init__`` is what lets a CUDA-first + priority list degrade to the PyTorch reference on a build without the + extension. Mirrors ``_require_cuda_activation`` in the activation ops. + """ + if not _EXT_AVAILABLE or _C is None: + raise RuntimeError(f"{what} requires the compiled rl_engine._C extension.") + missing = [name for name in names if not hasattr(_C, name)] + if missing: + raise RuntimeError( + f"{what} symbols ({', '.join(missing)}) are not compiled into _C. " + "Rebuild the extension with csrc/cuda/norm/rmsnorm.cu." + ) + + +#: Gate activations understood by the CUDA kernel, in binding order. ``swish`` is +#: an alias for ``silu``, as in vLLM's GDN block, which maps ``output_gate_type`` +#: "swish" to "silu" before constructing ``RMSNormGated``. +_GATE_ACTIVATIONS = {"silu": 0, "swish": 0, "sigmoid": 1} + + +def _check_gate_activation(activation: str) -> int: + if activation not in _GATE_ACTIVATIONS: + raise ValueError( + f"activation must be one of {sorted(_GATE_ACTIVATIONS)}, got {activation!r}" + ) + return _GATE_ACTIVATIONS[activation] + + +def _gate_activation_fp32(gate: torch.Tensor, activation: int) -> torch.Tensor: + """act(gate) in fp32, matching the kernel's ``gate_activation``.""" + gate32 = gate.float() + return torch.nn.functional.silu(gate32) if activation == 0 else torch.sigmoid(gate32) + + +def _gate_activation_grad_fp32(gate: torch.Tensor, activation: int) -> torch.Tensor: + """d act(gate) / d gate in fp32, matching ``gate_activation_grad``.""" + gate32 = gate.float() + sigma = torch.sigmoid(gate32) + if activation == 0: + return sigma * (1.0 + gate32 * (1.0 - sigma)) + return sigma * (1.0 - sigma) + + +class RMSNormGatedCuda(torch.autograd.Function): + """Autograd wrapper for the gated CUDA RMSNorm. + + Forward is the fused kernel. Backward is assembled from deterministic + pieces: ``dx`` from a row-local CUDA kernel, ``dweight`` from fp32 row + contributions reduced by the ascending-row left fold, and ``dgate`` purely + elementwise in fp32 (no reduction, so batch invariance is trivial). + """ + + @staticmethod + def forward(ctx, x, weight, gate, eps=1e-6, weight_offset=0.0, activation=0): + """ + Forward: + y = x * rsqrt(mean(x^2) + eps) * (weight_offset + weight) * act(gate) + + Input: + x, gate: [T, H], fp16/bf16/fp32 CUDA tensors of matching dtype + weight: [H] + activation: 0 = silu/swish, 1 = sigmoid + """ + assert x.is_cuda and weight.is_cuda and gate.is_cuda, "inputs must be CUDA tensors" + assert x.is_contiguous() and weight.is_contiguous() and gate.is_contiguous() + assert x.dim() == 2, "x must be [T, H]" + assert weight.dim() == 1, "weight must be [H]" + assert gate.shape == x.shape, "gate must match x" + assert _EXT_AVAILABLE and hasattr( + _C, "rmsnorm_gated_forward" + ), "Gated RMSNorm CUDA extension is unavailable. Rebuild with csrc/cuda/norm/rmsnorm.cu." + + y, rstd = _C.rmsnorm_gated_forward( + x, weight, gate, float(eps), float(weight_offset), int(activation) + ) + + ctx.save_for_backward(x, weight, gate, rstd) + ctx.eps = eps + ctx.weight_offset = float(weight_offset) + ctx.activation = int(activation) + + return y + + @staticmethod + def backward(ctx, grad_out): + x, weight, gate, rstd = ctx.saved_tensors + dy = grad_out.contiguous() + act = ctx.activation + + dx = _C.rmsnorm_gated_backward_dx(dy, x, weight, gate, rstd, ctx.weight_offset, act) + + # dweight: the gate is a per-element constant here, so the ungated row + # contributions apply once dy carries act(gate). + gate_act = _gate_activation_fp32(gate, act) + rows = rmsnorm_dweight_rows_fp32(x, dy.float() * gate_act, rstd=rstd) + dw = _fold_dweight_rows(rows, weight.dtype) + + # dgate: row-local and reduction-free. + normed = x.float() * rstd.unsqueeze(-1) + # Same guard as the kernels: an unconditional `+ 0.0` turns -0.0 weights into +0.0. + scale = weight.float() + if ctx.weight_offset != 0.0: + scale = scale + ctx.weight_offset + dgate = (dy.float() * normed * scale * _gate_activation_grad_fp32(gate, act)).to(gate.dtype) + + record_backward( + "rms_norm_gated", + kernel_id=( + "rl_engine._C.rmsnorm_gated_backward_dx" + "+rl_engine.ops.autograd.vjp_fp32.rmsnorm_dweight_rows_fp32" + "+rl_engine.ops.autograd.vjp_fp32.reduce_rows_fp32" + ), + impl="cuda_rmsnorm_gated_dx_declared_fp32_rowfold_dw", + family="cuda", + ) + + return dx, dw, dgate, None, None, None + + +def rmsnorm_gated_cuda(x, weight, gate, eps=1e-6, weight_offset=0.0, activation="silu"): + """ + use: + y = rmsnorm_gated_cuda(x, weight, gate) + y = rmsnorm_gated_cuda(x, weight, gate, activation="sigmoid") + """ + act = _check_gate_activation(activation) + return RMSNormGatedCuda.apply(x, weight, gate, eps, weight_offset, act) + + +class Qwen3NextRMSNormGatedCudaOp: + """CUDA gated RMSNorm for the Qwen3-Next GDN block. + + Deliberately not a subclass of :class:`RMSNormCudaOp`: it takes an extra + required tensor, so it cannot stand in for one. + + ``out = x * rstd * weight * silu(gate)``, every multiply in fp32 with a + single cast at the store. The weight is plain, not zero-centred, matching + vLLM's ``RMSNormGated`` with ``norm_before_gate=True`` and ``group_size=None`` + -- which is how the GDN block constructs it. The op has no ``norm_before_gate`` + or ``group_size`` parameter, so other configurations are not implemented. + + ``activation`` is fixed at construction; the registry constructs the + released config's ``"silu"``. + """ + + backward_impl = "cuda_rmsnorm_gated_dx_declared_fp32_rowfold_dw" + + #: The gated weight is plain; kept as an attribute so the surface matches + #: the ungated op and a zero-centred variant stays one subclass away. + weight_offset = 0.0 + + def __init__(self, activation: str = "silu") -> None: + _check_gate_activation(activation) + self.activation = activation + _require_cuda_symbols( + "Gated CUDA RMSNorm", "rmsnorm_gated_forward", "rmsnorm_gated_backward_dx" + ) + + def __call__(self, x, weight, gate, *, eps=1e-6): + return self.forward(x, weight, gate, eps=eps) + + def forward(self, x, weight, gate, *, eps=1e-6): + if gate.shape != x.shape: + raise ValueError(f"gate must match x, got {tuple(gate.shape)} vs {tuple(x.shape)}") + hidden = x.shape[-1] + x_2d = x.contiguous().view(-1, hidden) + gate_2d = gate.contiguous().view(-1, hidden) + y_2d = rmsnorm_gated_cuda( + x_2d, + weight.contiguous(), + gate_2d, + eps=eps, + weight_offset=self.weight_offset, + activation=self.activation, + ) + return y_2d.view_as(x) + + def parameter_vjp_contributions_fp32(self, *, x, weight, gate, grad_output, eps=1e-6): + if gate.shape != x.shape: + raise ValueError(f"gate must match x, got {tuple(gate.shape)} vs {tuple(x.shape)}") + hidden = x.shape[-1] + act = _GATE_ACTIVATIONS[self.activation] + _, rstd = _C.rmsnorm_gated_forward( + x.contiguous().reshape(-1, hidden), + weight.contiguous(), + gate.contiguous().reshape(-1, hidden), + float(eps), + float(self.weight_offset), + act, + ) + rows = rmsnorm_dweight_rows_fp32( + x, + grad_output.float() * _gate_activation_fp32(gate, act), + rstd=rstd.reshape(x.shape[:-1]), + ) return {"weight": rows} diff --git a/rl_engine/config/workload.py b/rl_engine/config/workload.py index 69b21c20f..79f93208a 100644 --- a/rl_engine/config/workload.py +++ b/rl_engine/config/workload.py @@ -264,10 +264,30 @@ def load_manifest(path: str | Path | None = None) -> WS1Manifest: raw = json.load(fh) if not isinstance(raw, dict): raise WorkloadError("manifest root must be a JSON object") - validate_manifest(raw) + if raw.get("scope") == "qwen3_next_norm_operators": + from rl_engine.validation.models.qwen3_next_workload import validate_norm_manifest + + validate_norm_manifest(raw) + else: + validate_manifest(raw) return WS1Manifest(raw=raw, path=manifest_path) +def workload_report(manifest: WS1Manifest) -> dict[str, Any]: + """The ``workload`` block the C3/C4 gate scripts attach to their JSON report. + + ``full_model_evidence`` is whatever the manifest declares, and ``None`` when + it declares nothing. + """ + return { + "workload_id": manifest.workload_id, + "scope": manifest.raw.get("scope", "qwen3_8b_dense"), + "fixture_identity_sha256": manifest.raw["fixture_identity_sha256"], + "model_id": manifest.model_identity["model_id"], + "full_model_evidence": manifest.raw.get("full_model_evidence"), + } + + def validate_manifest(raw: Mapping[str, Any]) -> None: """Hard-fail if any required C2 pin is missing or inconsistent.""" missing = [k for k in _REQUIRED_TOP_LEVEL if k not in raw] @@ -292,20 +312,25 @@ def validate_manifest(raw: Mapping[str, Any]) -> None: ) -def _validate_model_identity(identity: Mapping[str, Any]) -> None: +def _validate_model_identity( + identity: Mapping[str, Any], + *, + fingerprint: Mapping[str, Any] = _OFFICIAL_FINGERPRINT, + model_label: str = "Qwen3-8B Dense", +) -> None: for key in ("model_id", "revision", "config_fingerprint", "weight_snapshot"): if key not in identity: raise WorkloadError(f"model_identity missing {key!r}") fp = identity["config_fingerprint"] if not isinstance(fp, Mapping): raise WorkloadError("config_fingerprint must be an object") - for key, expected in _OFFICIAL_FINGERPRINT.items(): + for key, expected in fingerprint.items(): if key not in fp: raise WorkloadError(f"config_fingerprint missing {key!r}") if fp[key] != expected: raise WorkloadError( f"config_fingerprint {key}={fp[key]!r} does not match official " - f"Qwen3-8B Dense pin {expected!r}; architecture shrink is forbidden" + f"{model_label} pin {expected!r}; architecture shrink is forbidden" ) if not identity.get("exit_forbids_architecture_shrink", False): raise WorkloadError("exit_forbids_architecture_shrink must be true") diff --git a/rl_engine/reference/norm/qwen3_next_rms_norm.py b/rl_engine/reference/norm/qwen3_next_rms_norm.py new file mode 100644 index 000000000..14ecdcfa6 --- /dev/null +++ b/rl_engine/reference/norm/qwen3_next_rms_norm.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Qwen3-Next RMSNorm references (WS1 ground truth for RFC #428 C1). + +Qwen3-Next ships two RMSNorm conventions. They differ in how the weight is applied +and in where the dtype casts sit, so they are separate operators rather than a flag: + +``Qwen3NextRMSNorm`` (decoder / final norm) + Normalize in fp32, scale by ``(1 + weight)`` in fp32, cast once at the end. The + stored weight is zero-centred, so the ``1 +`` must be applied after the fp32 + upcast; folding it into a bf16 weight first rounds the offset. + +``Qwen3NextRMSNormGated`` (inside the Gated DeltaNet block) + Normalize in fp32, scale by a plain weight, then gate by ``silu(gate)``. + vLLM's ``RMSNormGated`` multiplies the weight in fp32; transformers casts the + normalized value back to the input dtype first. ``Qwen3NextRMSNormGatedOp`` + follows vLLM; ``Qwen3NextRMSNormGatedHFOp`` keeps the transformers convention + as a witness, and ``test_gated_conventions_diverge_in_low_precision`` pins that + the two differ in bf16 and agree bitwise in fp32. + +All three reuse :func:`shape_invariant_rstd`, a fixed-order reduction. They +reproduce the weight convention and cast order, not vLLM's reduction tree, and are +not bitwise equal to any vLLM path probed so far. Claim levels, measurements and +limitations are in ``docs/operators/qwen3-next-rms-norm.md`` and, for the gated +pair, ``docs/operators/qwen3-next-rms-norm-gated.md``. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.reference.norm.rms_norm import ( + NativeRMSNormOp, + check_norm_weight, + shape_invariant_rstd, +) + +__all__ = [ + "Qwen3NextRMSNormOp", + "Qwen3NextRMSNormGatedOp", + "Qwen3NextRMSNormGatedHFOp", +] + + +class Qwen3NextRMSNormOp(NativeRMSNormOp): + """Zero-centred RMSNorm: ``out = x * rstd * (1 + weight)``. + + Only the weight convention differs from :class:`NativeRMSNormOp`, so that is + all this overrides. The base applies the offset in fp32, after the upcast, + which is what ``transformers`` and vLLM both do -- folding ``1 +`` into a + bf16 weight beforehand would round the offset away. + """ + + weight_offset = 1.0 + + +class Qwen3NextRMSNormGatedOp: + """Gated RMSNorm used by the Gated DeltaNet block (vLLM/strict convention). + + ``out = (x * rstd * weight) * silu(gate)``, with every multiply in fp32 and + a single cast on the way out. This is the convention vLLM's ``RMSNormGated`` + uses with ``norm_before_gate=True``. Which gated convention is the strict + default is still open; see ``docs/operators/qwen3-next-rms-norm-gated.md``. + + Not a subclass of the plain op: it takes an extra tensor and its epilogue + differs, so it is not a drop-in substitute for one. + + The weight is plain, NOT zero-centred -- upstream initializes it to ones. + """ + + def __call__( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + return self.forward(x, weight, gate, eps=eps) + + def forward( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + return self._rms_norm_gated(x, weight, gate, eps=eps, output_dtype=x.dtype) + + def forward_fp32( + self, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float = 1e-6, + ) -> torch.Tensor: + """Ground-truth: fp32 in, fp32 out.""" + return self._rms_norm_gated(x, weight, gate, eps=eps, output_dtype=torch.float32) + + # ------------------------------------------------------------------ # + # Shared by both conventions; only `_scale_by_weight` differs. + # ------------------------------------------------------------------ # + @staticmethod + def _normalized( + x: torch.Tensor, weight: torch.Tensor, gate: torch.Tensor, *, eps: float + ) -> torch.Tensor: + check_norm_weight(x, weight) + if gate.shape != x.shape: + raise ValueError( + f"gate must match x, got tuple(gate.shape)={tuple(gate.shape)} " + f"vs tuple(x.shape)={tuple(x.shape)}" + ) + if gate.dtype != x.dtype or gate.device != x.device: + raise ValueError("gate must have the same dtype and device as x") + x_f = x.float() + rstd = shape_invariant_rstd(x_f, float(eps)).unsqueeze(-1) + return x_f * rstd + + @staticmethod + def _scale_by_weight( + normed: torch.Tensor, weight: torch.Tensor, input_dtype: torch.dtype + ) -> torch.Tensor: + """vLLM: the weight multiply stays in fp32.""" + del input_dtype + return normed * weight.float() + + @classmethod + def _rms_norm_gated( + cls, + x: torch.Tensor, + weight: torch.Tensor, + gate: torch.Tensor, + *, + eps: float, + output_dtype: torch.dtype, + ) -> torch.Tensor: + normed = cls._normalized(x, weight, gate, eps=eps) + scaled = cls._scale_by_weight(normed, weight, x.dtype) + # The gate activation is evaluated in fp32 and promotes the product. + gated = scaled * F.silu(gate.float()) + return gated.to(output_dtype) + + +class Qwen3NextRMSNormGatedHFOp(Qwen3NextRMSNormGatedOp): + """``transformers`` gated convention: weight multiply in the input dtype. + + Kept so the HF-vs-vLLM divergence documented in the module docstring stays + covered by a test rather than discovered in a drift report. Do NOT use this + for an L2 claim against vLLM rollout. + """ + + @staticmethod + def _scale_by_weight( + normed: torch.Tensor, weight: torch.Tensor, input_dtype: torch.dtype + ) -> torch.Tensor: + """transformers: round-trip through the input dtype first.""" + return weight * normed.to(input_dtype) diff --git a/rl_engine/reference/norm/rms_norm.py b/rl_engine/reference/norm/rms_norm.py index 76dc13b49..e08738a2e 100644 --- a/rl_engine/reference/norm/rms_norm.py +++ b/rl_engine/reference/norm/rms_norm.py @@ -86,6 +86,15 @@ def strict_add_rms_norm( return _strict_add_rms_norm(x, residual, weight, eps) +def check_norm_weight(x: torch.Tensor, weight: torch.Tensor) -> None: + """Shared shape guard for the RMSNorm family (plain, zero-centred, gated).""" + if weight.dim() != 1 or weight.shape[0] != x.shape[-1]: + raise ValueError( + f"weight must be 1-D of size x.shape[-1]={x.shape[-1]}, " + f"got tuple(weight.shape)={tuple(weight.shape)}" + ) + + def shape_invariant_rstd(x_f: torch.Tensor, eps: float) -> torch.Tensor: """Shape-invariant per-row rstd (the shared RMSNorm statistic). @@ -113,6 +122,10 @@ class NativeRMSNormOp: out = x * rsqrt(mean(x^2, dim=-1) + eps) * weight """ + #: Added to the weight in fp32 before it scales the normalized value. + #: Subclasses set 1.0 for the zero-centred ``(1 + w)`` convention. + weight_offset = 0.0 + def __init__(self) -> None: pass @@ -151,21 +164,24 @@ def forward_fp32( # ------------------------------------------------------------------ # # Helpers # ------------------------------------------------------------------ # - @staticmethod + @classmethod def _rms_norm( + cls, x: torch.Tensor, weight: torch.Tensor, *, eps: float, output_dtype: torch.dtype, ) -> torch.Tensor: - if weight.dim() != 1 or weight.shape[0] != x.shape[-1]: - raise ValueError( - f"weight must be 1-D of size x.shape[-1]={x.shape[-1]}, " - f"got tuple(weight.shape)={tuple(weight.shape)}" - ) + check_norm_weight(x, weight) x_f = x.float() rstd = shape_invariant_rstd(x_f, float(eps)).unsqueeze(-1) normed = x_f * rstd - out = normed * weight.float() + scale = weight.float() + # Guarded rather than unconditional: `0.0 + w` rewrites -0.0 to +0.0, + # which torch.equal does not notice but a bitwise comparison does. The + # plain path must stay bit-for-bit what it was. + if cls.weight_offset: + scale = cls.weight_offset + scale + out = normed * scale return out.to(output_dtype) diff --git a/rl_engine/runtime/registry.py b/rl_engine/runtime/registry.py index a190bbca1..9a6494d5d 100644 --- a/rl_engine/runtime/registry.py +++ b/rl_engine/runtime/registry.py @@ -159,6 +159,18 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): TRITON_RMS_NORM = "rl_engine.kernels.ops.triton.rmsnorm_triton.RMSNormTritonOp" PYTORCH_NATIVE_RMS_NORM = "rl_engine.kernels.ops.pytorch.norm.rms_norm.NativeRMSNormOp" + # Zero-centred RMSNorm (Qwen3-Next decoder / final norm) + CUDA_QWEN3_NEXT_RMS_NORM = "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp" + PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM = ( + "rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp" + ) + + # Gated RMSNorm (Qwen3-Next Gated DeltaNet block) + CUDA_RMS_NORM_GATED = "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp" + PYTORCH_NATIVE_RMS_NORM_GATED = ( + "rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp" + ) + # Generic fallback TRITON_GENERIC = "rl_engine.kernels.ops.triton.generic.TritonOp" PYTORCH_ATTN = "rl_engine.kernels.ops.pytorch.attention.NativeAttentionOp" @@ -594,6 +606,14 @@ def __init__(self): OpBackend.CUDA_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ], + "rms_norm_gated": [ + OpBackend.CUDA_RMS_NORM_GATED, + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + ], + "qwen3_next_rms_norm": [ + OpBackend.CUDA_QWEN3_NEXT_RMS_NORM, + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + ], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [ @@ -649,6 +669,8 @@ def __init__(self): ], "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.TRITON_SILU, OpBackend.PYTORCH_NATIVE_SILU], @@ -694,6 +716,8 @@ def __init__(self): OpBackend.TRITON_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [ OpBackend.TRITON_EMBEDDING, @@ -722,6 +746,8 @@ def __init__(self): "batch_invariant_logp": [OpBackend.PYTORCH_BATCH_INVARIANT_LOGP], "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], + "rms_norm_gated": [OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED], + "qwen3_next_rms_norm": [OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.PYTORCH_NATIVE_SILU], @@ -762,6 +788,12 @@ def __init__(self): OpBackend.ASCEND_RMS_NORM, OpBackend.PYTORCH_NATIVE_RMS_NORM, ] + self._priority_map["npu"]["rms_norm_gated"] = [ + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + ] + self._priority_map["npu"]["qwen3_next_rms_norm"] = [ + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + ] self._priority_map["npu"]["embedding"] = [ OpBackend.ASCEND_EMBEDDING, OpBackend.PYTORCH_NATIVE_EMBEDDING, diff --git a/rl_engine/validation/models/qwen3_next_norm_manifest.json b/rl_engine/validation/models/qwen3_next_norm_manifest.json new file mode 100644 index 000000000..9cf5977ef --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_norm_manifest.json @@ -0,0 +1,759 @@ +{ + "version": "1.0", + "workload_id": "qwen3-next-80b-a3b-norm-c3-c4-v1", + "seed": 20260812, + "model_identity": { + "model_id": "Qwen/Qwen3-Next-80B-A3B-Instruct", + "revision": "9c7f2fbe84465e40164a94cc16cd30b6999b0cc7", + "config_fingerprint": { + "num_hidden_layers": 48, + "hidden_size": 2048, + "intermediate_size": 5120, + "num_attention_heads": 16, + "num_key_value_heads": 2, + "head_dim": 256, + "vocab_size": 151936, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "num_experts": 512, + "num_experts_per_tok": 10 + }, + "exit_forbids_architecture_shrink": true, + "weight_snapshot": { + "pin_method": "revision-and-sha256", + "total_size_bytes": 162649725440, + "index_file": "model.safetensors.index.json", + "index_sha256": "0e5274f538d750abdd7fd55a36bad64100f9e881d59c65f6a59f364d2f8b6dd7", + "content_hash_algorithm": "sha256-of-sorted-shard-records-v1", + "content_hash": "00a6a3b5573ccbac2e01fadb673c2cd7bca172d40ddd73f95ed7afa5cbdb9dcd", + "shards": [ + { + "filename": "model-00001-of-00041.safetensors", + "size_bytes": 3999619256, + "sha256": "d8908fd48650169854b5ed815cb01bfc1741152cc58795ec99188862c6a29e11" + }, + { + "filename": "model-00002-of-00041.safetensors", + "size_bytes": 3999841784, + "sha256": "c753c9bfaca220781d4030c3a99e69b4a256434c9d6ec223f5147edc265289df" + }, + { + "filename": "model-00003-of-00041.safetensors", + "size_bytes": 3999515584, + "sha256": "51aaa14dd50c5ab90c363227bfb1ac51182118f588e0141dcaadda012548407c" + }, + { + "filename": "model-00004-of-00041.safetensors", + "size_bytes": 3999842000, + "sha256": "82a33096134fb6e7751a423f142e79b1bbe89f45242b75397c2d7170fffa75bd" + }, + { + "filename": "model-00005-of-00041.safetensors", + "size_bytes": 3999842208, + "sha256": "41abda7bf93f27c36cb28382114f5b33defa6dbfd154abb6886a6fec20f9e479" + }, + { + "filename": "model-00006-of-00041.safetensors", + "size_bytes": 3999853216, + "sha256": "a7794475040ebd62c9a1f9c94c17e7c600873e14fcabc656b32faea09bc7fd2d" + }, + { + "filename": "model-00007-of-00041.safetensors", + "size_bytes": 3999841912, + "sha256": "64d7e90d00ce15cc8bbf7677ad07bf3cffc9e90a1b777e3334fa15ffea219d6b" + }, + { + "filename": "model-00008-of-00041.safetensors", + "size_bytes": 3999842000, + "sha256": "9ace9b99e71490d619656956d457744f727647b3f36fa3d801798a5156599d35" + }, + { + "filename": "model-00009-of-00041.safetensors", + "size_bytes": 3999843192, + "sha256": "5b1b374c9e6100077c446a293566c1f644475c94739ab08b7fa26ba847110216" + }, + { + "filename": "model-00010-of-00041.safetensors", + "size_bytes": 3999517808, + "sha256": "38e51dc850a39c325324e6ddd23347e6e3e76cba1341aeef2e1f260ca2cb9f49" + }, + { + "filename": "model-00011-of-00041.safetensors", + "size_bytes": 4000181296, + "sha256": "c06719dc79fbbc796b8751ccbb29c8ecbdb686de9f89e3c741b5bbfec203cf0c" + }, + { + "filename": "model-00012-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "349036c3e3f6a8645af47906dae7272bd5352f822a3c354a16ef4084b3555679" + }, + { + "filename": "model-00013-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "bfab17d1535ef23e99e54ff3156ce1456aa5f0a4a93b281a7adb9d8b921dc829" + }, + { + "filename": "model-00014-of-00041.safetensors", + "size_bytes": 3999843984, + "sha256": "318d10df78f1647189941884acc01616a6a78d5d45d74ea30955be7f509c580a" + }, + { + "filename": "model-00015-of-00041.safetensors", + "size_bytes": 4000181736, + "sha256": "aa3400abde789ecca625b8cea37aa31232bf6696923f10981d028518a2393173" + }, + { + "filename": "model-00016-of-00041.safetensors", + "size_bytes": 3999517256, + "sha256": "3776700f68222f0174bb14fea8480da41d928a1bfd6b5317c4fc2b14006d506e" + }, + { + "filename": "model-00017-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "052726cd55224d3d217dac93eaed060215b8029afb44861c0d927d87c6a046ab" + }, + { + "filename": "model-00018-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "dd1b8d843e0013ed3beb6ba7fc36dc4b1f2df5d71b40520f82f369cd8fce5cba" + }, + { + "filename": "model-00019-of-00041.safetensors", + "size_bytes": 3999844096, + "sha256": "bfe5489c61334f9ce9bc6b369c3edf44420d7f585c6438f3afb98a120b26e6f3" + }, + { + "filename": "model-00020-of-00041.safetensors", + "size_bytes": 3999855040, + "sha256": "eba3113cbd34304751e148502f314be5ceae6e7d830bd33f2c0cf3bb8dfb28c2" + }, + { + "filename": "model-00021-of-00041.safetensors", + "size_bytes": 3999843792, + "sha256": "9ddd72e8ada4480e86e4230639210df1039802db9f78537673c19117d2e69b95" + }, + { + "filename": "model-00022-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "c6fc2bb1762d34b8db83fb11566e462c9166e132bb2bfcef45a33a4bd41c02db" + }, + { + "filename": "model-00023-of-00041.safetensors", + "size_bytes": 3999517464, + "sha256": "484bf8c327f7f7aaa20fbc4f6800d4aab7870749335b6f6bbfdb83570a1f66e2" + }, + { + "filename": "model-00024-of-00041.safetensors", + "size_bytes": 3999844264, + "sha256": "cfbb94709f5dac71ffe144ee3b9a976a49a2a9b9863caa29620fbc65820c2342" + }, + { + "filename": "model-00025-of-00041.safetensors", + "size_bytes": 4000181296, + "sha256": "b5799d19dcccfb17b108b5349900a4c42eb7bc241149cb7d1b38b48c707070cf" + }, + { + "filename": "model-00026-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "98003819929196d2f6f05c42603e9124a7a7853f372089eed3ffbd0c244fbdac" + }, + { + "filename": "model-00027-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "3d6b8b9d1c0c262cc8a63c41c3f75053ab8b31b6107130800718048da64b86db" + }, + { + "filename": "model-00028-of-00041.safetensors", + "size_bytes": 3999843984, + "sha256": "836d4a85cc72243c6caee2b6cf470cd598d6ab817e0685daf1128fd97bb133e3" + }, + { + "filename": "model-00029-of-00041.safetensors", + "size_bytes": 3999855320, + "sha256": "7b083d02c732458cfaaee8b0d1407626ad5ee10fc7eeed6b294b7d6853ac30e8" + }, + { + "filename": "model-00030-of-00041.safetensors", + "size_bytes": 3999843672, + "sha256": "e81006eeaa1152c67798640be46b710e87d44a307effca60cbdb7e9d1a3b268b" + }, + { + "filename": "model-00031-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "4fd46dbc55ceefa6146a7687cd02d040087c58b0f7ea42c4cd7dc66b838f531c" + }, + { + "filename": "model-00032-of-00041.safetensors", + "size_bytes": 3999843880, + "sha256": "7edf6738b9bc204f0b39e9c7899515d1a1ed9d7d5f7c1cce32c981120c4f034d" + }, + { + "filename": "model-00033-of-00041.safetensors", + "size_bytes": 3999517688, + "sha256": "9f0f735e75e8d648b4b8d601a598d489fc4e4c7876da4f3406bf3bcc8990a181" + }, + { + "filename": "model-00034-of-00041.safetensors", + "size_bytes": 4000181496, + "sha256": "3499ba70dbb6a9acf87509c3b82c5790461a73f8d0d12f0e3699354f34dc02b7" + }, + { + "filename": "model-00035-of-00041.safetensors", + "size_bytes": 3999843792, + "sha256": "de04b5e39b9494b44becacf22fbe9c02890ef323a3dfe37640a2bfa80335504c" + }, + { + "filename": "model-00036-of-00041.safetensors", + "size_bytes": 3999517472, + "sha256": "3eddb88c62194513f94ef3b1c11577ad324fdcc21265ff7b11a4b50d50b03c5f" + }, + { + "filename": "model-00037-of-00041.safetensors", + "size_bytes": 3999843872, + "sha256": "01992493e00d45d5fec4b5989392f401c9ebe3e13973edaaf1973577e132ca11" + }, + { + "filename": "model-00038-of-00041.safetensors", + "size_bytes": 3999844264, + "sha256": "e328c6d2e5f2a42eb99a0335bde7aa4c42cbd72b88fb988d5a2c96ba6231e95e" + }, + { + "filename": "model-00039-of-00041.safetensors", + "size_bytes": 3999854888, + "sha256": "15063cd0a46af3fa37e74590534f62a81f1ede5340e7640713bf98f39f20961d" + }, + { + "filename": "model-00040-of-00041.safetensors", + "size_bytes": 3365572496, + "sha256": "e458c4f4345d7082e7063f381ee1962725f84c44d53a5b70afb77e601ff22017" + }, + { + "filename": "model-00041-of-00041.safetensors", + "size_bytes": 3301131296, + "sha256": "1f3ac4d828f7e08dd14eb4dc6f282139ae1fa0d894e43f247da2539d2ef43826" + } + ], + "weight_files_total_size_bytes": 162659161528 + } + }, + "chain_semantics": { + "execution_dtype": "bfloat16", + "reference_dtype": "float32", + "accumulation_dtype": "float32", + "temperature": 1.0, + "loss_reduction": "sum_over_active_tokens_then_optional_mean_by_active_count", + "logprob_selection": "selected_token_logprob_on_active_mask", + "active_token_policy": "active selected tokens only", + "aggregates": [ + "max_abs_dlogp", + "approx_kl0", + "clipfrac0" + ], + "clip_interval": [ + 0.8, + 1.2 + ], + "clip_interval_note": "Pinned for clipfrac0; must match C1 chain_logprob_aggregates.default_clip_interval unless an explicit contract revision changes both.", + "comparison_roles_source": "rl_engine/kernels/gtest/tolerance_contract.json", + "forbidden_comparison_roles": [ + "baseline", + "singleton_aggregate" + ], + "singleton_aggregate_note": "singleton_aggregate is a C2 execution/aggregation mode only. It must never populate comparison_lhs_role or comparison_rhs_role.", + "tf32_policy_ref": "rl_engine/kernels/gtest/tolerance_contract.json#/policy/tf32", + "tf32_note": "WS1 TF32 enable/disable is owned by the C1 contract; C2 gates must not introduce a private TF32 policy.", + "report_naming": { + "comparison_lhs_role": "from_c1_by_report_kind", + "comparison_rhs_role": "from_c1_by_report_kind", + "forbidden_in_reports": [ + "baseline", + "singleton_aggregate" + ], + "singleton_aggregate_is": "c2_execution_aggregation_mode_only", + "note": "C2 freezes naming rules; C3+ emit reports that must obey these roles." + }, + "backend_actual_semantics": { + "c2_representative_actual_source": "scripts/ws1_candidate_evidence.py runtime execution", + "full_model_runtime_observed_actual_owner": [ + "C3", + "C8", + "C10", + "C11" + ], + "note": "C2 executes every representative case and records runtime-observed actual backend/kernel provenance. Later children own full-model dispatch provenance." + } + }, + "stochastic_policy": { + "dropout": 0.0, + "attention_dropout": 0.0, + "sampling_in_logprob_parity": false, + "canonical_gate_uses_dropout_zero": true, + "rng_source": "manifest_seed_plus_logical_sample_token_identity", + "undeclared_randomness": "hard_fail", + "retained_stochastic_ops": [] + }, + "primary_matrix": { + "description": "Fixed #150 Batch \u00d7 Chunked-Prefill matrix prerequisite workload cells.", + "N": 4, + "batch_size_bn": 4, + "sample_ids": [ + "s0", + "s1", + "s2", + "s3" + ], + "sample_order_fixed": true, + "batch_permutation": { + "enabled": true, + "permutation": [ + 2, + 0, + 3, + 1 + ], + "target_sample_position_in_bn": 0, + "note": "Permutation exercises layout invariance; logical compare restores sample_id order." + }, + "chunk": { + "chunk_size_tokens": 7, + "require_ge_2_chunks": true, + "non_divisible_case": true, + "note": "Longest primary seq_len=19 with chunk_size=7 yields chunks [7,7,5]." + }, + "cells": [ + { + "cell_id": "B1-singleton_aggregate/full", + "batch_mode": "singleton_aggregate", + "batch_size_per_run": 1, + "num_runs": 4, + "prefill_mode": "full", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "BN/full", + "batch_mode": "batched", + "batch_size_per_run": 4, + "num_runs": 1, + "prefill_mode": "full", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "B1-singleton_aggregate/chunked", + "batch_mode": "singleton_aggregate", + "batch_size_per_run": 1, + "num_runs": 4, + "prefill_mode": "chunked", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + }, + { + "cell_id": "BN/chunked", + "batch_mode": "batched", + "batch_size_per_run": 4, + "num_runs": 1, + "prefill_mode": "chunked", + "aggregation": { + "order": "sample_ids", + "denominator": "active_token_count_across_all_samples" + } + } + ] + }, + "fixtures": { + "prompt_template": "ws1_fixed_token_fixture", + "dtype_for_token_tensors": "int64", + "position_ids": { + "basis": "logical_zero_based_per_sample", + "reset_after_pack_boundary": true + }, + "attention_mask": { + "active_value": 1, + "padding_value": 0, + "causal": true + }, + "primary_seq_len": 19, + "primary_prompt_len": 8, + "short_seq_len": 8, + "long_seq_len": 32, + "varlen_seq_lens": [ + 11, + 16, + 13, + 19 + ], + "padding": { + "modes": [ + "right", + "left" + ], + "pad_token_id": 151643, + "primary_padded_len": 20 + }, + "packing": { + "status": "supported", + "implementation": "rl_engine.kernels.ops.pytorch.packing.pack.NativePackOp", + "packed_fixture": { + "sample_order": [ + "s0", + "s1", + "s2", + "s3" + ], + "segment_lengths": [ + 11, + 16, + 13, + 19 + ], + "total_tokens": 59, + "restore_key": [ + "sample_id", + "token_position" + ] + } + }, + "loss_mask": { + "prompt_tokens_active": false, + "completion_tokens_active": true + }, + "samples": [ + { + "sample_id": "s0", + "seq_len": 11, + "prompt_len": 8, + "token_ids": [ + 100, + 101, + 102, + 103, + 104, + 105, + 106, + 107, + 200, + 201, + 202 + ] + }, + { + "sample_id": "s1", + "seq_len": 16, + "prompt_len": 8, + "token_ids": [ + 110, + 111, + 112, + 113, + 114, + 115, + 116, + 117, + 210, + 211, + 212, + 213, + 214, + 215, + 216, + 217 + ] + }, + { + "sample_id": "s2", + "seq_len": 13, + "prompt_len": 8, + "token_ids": [ + 120, + 121, + 122, + 123, + 124, + 125, + 126, + 127, + 220, + 221, + 222, + 223, + 224 + ] + }, + { + "sample_id": "s3", + "seq_len": 19, + "prompt_len": 8, + "token_ids": [ + 130, + 131, + 132, + 133, + 134, + 135, + 136, + 137, + 230, + 231, + 232, + 233, + 234, + 235, + 236, + 237, + 238, + 239, + 240 + ] + } + ], + "short_full_model_fixture": { + "fixture_id": "short_full_model_seq8", + "seq_len": 8, + "prompt_len": 4, + "token_ids": [ + 310, + 311, + 312, + 313, + 410, + 411, + 412, + 413 + ], + "note": "Shorter sequence on full architecture+weights only; never shrinks layers/hidden/heads/vocab.", + "candidate_case_ids": [ + "short_full_model_seq8_qwen3_next_rms_norm", + "short_full_model_seq8_rms_norm_gated" + ] + }, + "long_full_model_fixture": { + "fixture_id": "long_full_model_seq32", + "seq_len": 32, + "prompt_len": 16, + "token_ids": [ + 500, + 501, + 502, + 503, + 504, + 505, + 506, + 507, + 508, + 509, + 510, + 511, + 512, + 513, + 514, + 515, + 600, + 601, + 602, + 603, + 604, + 605, + 606, + 607, + 608, + 609, + 610, + 611, + 612, + 613, + 614, + 615 + ], + "note": "Long fixed sequence on the same full architecture and pinned weight snapshot.", + "candidate_case_ids": [ + "long_full_model_seq32_qwen3_next_rms_norm", + "long_full_model_seq32_rms_norm_gated" + ] + }, + "representative_full_model_fixture": { + "fixture_id": "rep_full_model_seq16", + "seq_len": 16, + "prompt_len": 8, + "sample_ids": [ + "s0", + "s1", + "s2", + "s3" + ], + "note": "Primary variable-length matrix fixture; full architecture+weights.", + "candidate_case_ids": [ + "rep_full_model_seq16_qwen3_next_rms_norm", + "rep_full_model_seq16_rms_norm_gated" + ] + }, + "prompt_lens": [ + 8, + 8, + 8, + 8 + ], + "completion_lens": [ + 3, + 8, + 5, + 11 + ], + "max_completion_len": 11 + }, + "logical_identity": { + "key": [ + "sample_id", + "token_position" + ], + "token_position_basis": "logical_unpadded_index_in_sample", + "restore_before_compare_after": [ + "pad", + "pack", + "chunk", + "batch_permute" + ], + "gradient_singleton_aggregate": { + "definition": "N independent B=1 runs of the same N logical samples, aggregated with fixed sample order and active-token denominator", + "compare_to": "single B=N run of the same logical sample/token multiset", + "forbid_different_sample_sets": true + } + }, + "capabilities": { + "required_chain_ops": [ + { + "op": "qwen3_next_rms_norm", + "status": "required" + }, + { + "op": "rms_norm_gated", + "status": "required" + } + ], + "operator_spec_map": { + "qwen3_next_rms_norm": "qwen3_next_rms_norm", + "rms_norm_gated": "rms_norm_gated" + } + }, + "backend_profiles": { + "cuda_bf16": { + "backend_family": "cuda", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "declared", + "expected_backend_id": "cuda", + "expected_kernel_config_id": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp", + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "declared", + "expected_backend_id": "cuda", + "expected_kernel_config_id": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp", + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + }, + "triton_cuda_bf16": { + "backend_family": "triton", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + }, + "ascend_bf16": { + "backend_family": "ascend", + "execution_dtype": "bfloat16", + "required_nodes": [ + { + "node": "qwen3_next_rms_norm", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + }, + { + "node": "rms_norm_gated", + "status": "missing_required", + "expected_backend_id": null, + "expected_kernel_config_id": null, + "algorithm_property": "fixed row reduction and ordered FP32 parameter fold" + } + ] + } + }, + "representative_cases": [ + { + "case_id": "short_full_model_seq8_qwen3_next_rms_norm", + "fixture_id": "short_full_model_seq8", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "short_full_model_seq8_rms_norm_gated", + "fixture_id": "short_full_model_seq8", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "long_full_model_seq32_qwen3_next_rms_norm", + "fixture_id": "long_full_model_seq32", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "long_full_model_seq32_rms_norm_gated", + "fixture_id": "long_full_model_seq32", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "rep_full_model_seq16_qwen3_next_rms_norm", + "fixture_id": "rep_full_model_seq16", + "operator_spec": "qwen3_next_rms_norm", + "hidden": 2048, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + }, + { + "case_id": "rep_full_model_seq16_rms_norm_gated", + "fixture_id": "rep_full_model_seq16", + "operator_spec": "rms_norm_gated", + "hidden": 128, + "architecture_identity": "qwen3_next_80b_a3b_norm_operators" + } + ], + "fixture_identity_sha256": "3ac9490ffbb1486edac6ea509811fd4d5affbe880b208bd7d9b052b0ce418442", + "provenance_boundary": { + "scope": "Synthetic norm inputs at checkpoint dimensions; not checkpoint execution or model-level L2.", + "runtime_verified": false + }, + "scope": "qwen3_next_norm_operators", + "full_model_evidence": false +} diff --git a/rl_engine/validation/models/qwen3_next_workload.py b/rl_engine/validation/models/qwen3_next_workload.py new file mode 100644 index 000000000..ee56f90da --- /dev/null +++ b/rl_engine/validation/models/qwen3_next_workload.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Qwen3-Next norm-only C3/C4 workload, explicitly separate from the Dense chain.""" + +from collections.abc import Mapping +from typing import Any + +from rl_engine.config.workload import ( + WorkloadError, + _validate_backend_profiles, + _validate_fixtures, + _validate_logical_identity, + _validate_model_identity, + _validate_primary_matrix, + _validate_stochastic_policy, + manifest_identity_hash, +) + +MODEL_ID = "Qwen/Qwen3-Next-80B-A3B-Instruct" +REVISION = "9c7f2fbe84465e40164a94cc16cd30b6999b0cc7" +FINGERPRINT = { + "num_hidden_layers": 48, + "hidden_size": 2048, + "intermediate_size": 5120, + "num_attention_heads": 16, + "num_key_value_heads": 2, + "head_dim": 256, + "vocab_size": 151936, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "num_experts": 512, + "num_experts_per_tok": 10, +} +NORM_OPS = {"qwen3_next_rms_norm": 2048, "rms_norm_gated": 128} + + +def validate_norm_manifest(raw: Mapping[str, Any]) -> None: + identity = raw["model_identity"] + if identity["model_id"] != MODEL_ID or identity["revision"] != REVISION: + raise WorkloadError("Qwen3-Next norm workload requires the pinned official checkpoint") + _validate_model_identity(identity, fingerprint=FINGERPRINT, model_label="Qwen3-Next") + if ( + raw.get("scope") != "qwen3_next_norm_operators" + or raw.get("full_model_evidence") is not False + ): + raise WorkloadError("norm workload cannot claim full-model evidence") + _validate_stochastic_policy(raw["stochastic_policy"]) + _validate_primary_matrix(raw["primary_matrix"], raw["fixtures"]) + _validate_fixtures(raw["fixtures"], raw["primary_matrix"]) + _validate_logical_identity(raw["logical_identity"]) + caps = raw["capabilities"] + if {e["op"] for e in caps["required_chain_ops"]} != set(NORM_OPS): + raise WorkloadError("norm workload must contain exactly the two Qwen3-Next norms") + if any(e["status"] != "required" for e in caps["required_chain_ops"]): + raise WorkloadError("both norm operators are required") + _validate_backend_profiles(raw["backend_profiles"], caps) + cases = {case["case_id"]: case for case in raw["representative_cases"]} + if len(cases) != len(raw["representative_cases"]): + raise WorkloadError("duplicate norm case ID") + referenced = set() + for key in ( + "short_full_model_fixture", + "long_full_model_fixture", + "representative_full_model_fixture", + ): + fixture = raw["fixtures"][key] + for case_id in fixture["candidate_case_ids"]: + if case_id not in cases or cases[case_id]["fixture_id"] != fixture["fixture_id"]: + raise WorkloadError("norm case fixture binding mismatch") + referenced.add(case_id) + if referenced != set(cases): + raise WorkloadError("unreferenced norm case") + if {case["operator_spec"] for case in cases.values()} != set(NORM_OPS): + raise WorkloadError("representative cases must cover both norms") + for case in cases.values(): + if case["hidden"] != NORM_OPS[case["operator_spec"]]: + raise WorkloadError("norm case hidden dimension does not match checkpoint") + if case["architecture_identity"] != "qwen3_next_80b_a3b_norm_operators": + raise WorkloadError("norm cases cannot claim Dense or full-model architecture evidence") + if raw["fixture_identity_sha256"] != manifest_identity_hash(raw): + raise WorkloadError("Qwen3-Next fixture identity hash mismatch") + + +def validate_norm_dimensions(raw: Mapping[str, Any], op: str, hidden: int, head_dim: int) -> None: + if raw.get("scope") != "qwen3_next_norm_operators": + return + if op not in NORM_OPS or hidden != 2048 or head_dim != 128: + raise WorkloadError( + "Qwen3-Next norm gate requires its two norm ops, --hidden 2048 --head-dim 128" + ) diff --git a/rl_engine/validation/operators/gradient_adapters.py b/rl_engine/validation/operators/gradient_adapters.py index f22d80285..9831fb670 100644 --- a/rl_engine/validation/operators/gradient_adapters.py +++ b/rl_engine/validation/operators/gradient_adapters.py @@ -53,6 +53,7 @@ class GradientAdapterSpec: source_files: tuple[str, ...] shape_dependent_bwd_accum: str = "forbidden" atomic_add: str = "forbidden" + model_id: str | None = None @dataclass(frozen=True) @@ -115,6 +116,26 @@ def to_dict(self) -> dict[str, Any]: "csrc/cuda/norm/rmsnorm.cu", ), ), + "qwen3_next_rms_norm": GradientAdapterSpec( + op_name="qwen3_next_rms_norm", + chain_node="qwen3_next_rms_norm", + op_class="reduction", + spec_name="qwen3_next_rms_norm", + tensors=(_DX, _DWEIGHT), + requirement="required", + source_files=("rl_engine/backends/cuda/norm/rmsnorm.py", "csrc/cuda/norm/rmsnorm.cu"), + model_id="Qwen/Qwen3-Next-80B-A3B-Instruct", + ), + "rms_norm_gated": GradientAdapterSpec( + op_name="rms_norm_gated", + chain_node="rms_norm_gated", + op_class="reduction", + spec_name="rms_norm_gated", + tensors=(_DX, _DWEIGHT, _DGATE), + requirement="required", + source_files=("rl_engine/backends/cuda/norm/rmsnorm.py", "csrc/cuda/norm/rmsnorm.cu"), + model_id="Qwen/Qwen3-Next-80B-A3B-Instruct", + ), "qk_norm": GradientAdapterSpec( op_name="qk_norm", chain_node="qk_norm", @@ -291,18 +312,27 @@ def get_adapter(op_name: str) -> GradientAdapterSpec: raise KeyError(f"unknown gradient adapter {op_name!r}") from exc -def required_gradient_adapters() -> tuple[GradientAdapterSpec, ...]: +def required_gradient_adapters( + manifest: WS1Manifest | None = None, +) -> tuple[GradientAdapterSpec, ...]: + selected = manifest or load_manifest() + model_id = selected.model_identity["model_id"] + subset = selected.raw.get("scope") == "qwen3_next_norm_operators" return tuple( spec for spec in GRADIENT_ADAPTERS.values() if spec.requirement in ("required", "layout_supported") + and spec.model_id in (None, model_id) + and (not subset or spec.model_id == model_id) ) -def required_forward_adapters() -> tuple[GradientAdapterSpec, ...]: +def required_forward_adapters( + manifest: WS1Manifest | None = None, +) -> tuple[GradientAdapterSpec, ...]: """Same enumerable WS1 ops as C4; C3 reuses the registry, not a second list.""" - return required_gradient_adapters() + return required_gradient_adapters(manifest) @dataclass(frozen=True) @@ -620,9 +650,9 @@ def _row_parameters( head_dim: int = 16, ) -> dict[str, torch.Tensor]: """Config-independent trainable parameters, built in the execution dtype.""" - if op_name == "rms_norm": + if op_name in {"rms_norm", "qwen3_next_rms_norm"}: return {"weight": _shared_parameter((hidden,), device=device, dtype=dtype, offset=1)} - if op_name == "qk_norm": + if op_name in {"qk_norm", "rms_norm_gated"}: return {"weight": _shared_parameter((head_dim,), device=device, dtype=dtype, offset=1)} if op_name == "det_gemm": return {"b": _shared_parameter((hidden, hidden), device=device, dtype=dtype, offset=2)} @@ -663,12 +693,19 @@ def _row_inputs( """ n = len(keys) leading = (n,) - if op_name == "rms_norm": + if op_name in {"rms_norm", "qwen3_next_rms_norm"}: return { "x": _stack_rows(keys, leading, (hidden,), device=device, dtype=dtype), "weight": params["weight"], "eps": 1.0e-6, } + if op_name == "rms_norm_gated": + return { + "x": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype), + "gate": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype, offset=7), + "weight": params["weight"], + "eps": 1.0e-6, + } if op_name == "qk_norm": return { "x": _stack_rows(keys, leading, (head_dim,), device=device, dtype=dtype), @@ -1228,6 +1265,15 @@ def resolve_profile_candidate( manifest: WS1Manifest | None = None, ) -> dict[str, Any]: m = manifest if manifest is not None else load_manifest() + subset = m.raw.get("scope") == "qwen3_next_norm_operators" + if adapter.model_id not in (None, m.model_identity["model_id"]) or ( + subset and adapter.model_id != m.model_identity["model_id"] + ): + return { + "status": "absent_not_required", + "expected_backend_id": None, + "candidate_path": None, + } if adapter.requirement == "absent_not_required": return { "status": "absent_not_required", diff --git a/rl_engine/validation/operators/operator_inputs.py b/rl_engine/validation/operators/operator_inputs.py index ca3b7c120..1cd03d31f 100644 --- a/rl_engine/validation/operators/operator_inputs.py +++ b/rl_engine/validation/operators/operator_inputs.py @@ -26,6 +26,8 @@ def make_operator_inputs( ) -> dict[str, Any]: builders = { "rms_norm": _make_rms_norm_inputs, + "rms_norm_gated": _make_rms_norm_gated_inputs, + "qwen3_next_rms_norm": _make_rms_norm_inputs, "qk_norm": _make_qk_norm_inputs, "pack": _make_pack_inputs, "matmul": _make_matmul_inputs, @@ -54,6 +56,8 @@ def operator_shape_name(op_name: str, args: argparse.Namespace) -> str: vocab = _arg_int(args, "vocab", DEFAULT_VOCAB) names = { "rms_norm": f"{batch}x{seq}x{_normalized_dim(args)}", + "rms_norm_gated": f"{batch}x{seq}x{_arg_int(args, 'head_dim', DEFAULT_HEAD_DIM)}", + "qwen3_next_rms_norm": f"{batch}x{seq}x{_normalized_dim(args)}", "qk_norm": f"{batch}x{seq}x{_arg_int(args, 'n_heads', DEFAULT_N_HEADS)}x" f"{_arg_int(args, 'head_dim', DEFAULT_HEAD_DIM)}", "pack": f"{batch}x{seq}x{_normalized_dim(args)}", @@ -91,6 +95,25 @@ def _make_rms_norm_inputs( } +def _make_rms_norm_gated_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + """Gated RMSNorm inputs at the GDN block's width. + + The gated norm normalizes over ``linear_value_head_dim`` (128 for + Qwen3-Next), not ``hidden_size``, so it follows ``head_dim`` rather than + ``normalized_dim``. + """ + batch, seq = _batch_seq(args) + head_dim = _arg_int(args, "head_dim", DEFAULT_HEAD_DIM) + return { + "x": _floating_tensor((batch, seq, head_dim), args, dtype, device, offset=0), + "weight": _floating_tensor((head_dim,), args, dtype, device, offset=1), + "gate": _floating_tensor((batch, seq, head_dim), args, dtype, device, offset=2), + "eps": _arg_float(args, "eps", DEFAULT_RMS_EPS), + } + + def _make_qk_norm_inputs( args: argparse.Namespace, dtype: torch.dtype, device: torch.device ) -> dict[str, Any]: diff --git a/rl_engine/validation/operators/operator_specs.py b/rl_engine/validation/operators/operator_specs.py index 2323f4787..7f17c369d 100644 --- a/rl_engine/validation/operators/operator_specs.py +++ b/rl_engine/validation/operators/operator_specs.py @@ -46,6 +46,30 @@ def _load_object(path: str) -> Any: }, grad_input_names=("x", "weight"), ), + "rms_norm_gated": OperatorSpec( + name="rms_norm_gated", + op_class="reduction", + gold_path=("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp"), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormGatedOp"), + "cuda": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp"), + "cuda-sm90": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormGatedCudaOp"), + }, + grad_input_names=("x", "weight", "gate"), + ), + "qwen3_next_rms_norm": OperatorSpec( + name="qwen3_next_rms_norm", + op_class="reduction", + gold_path=("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp"), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.norm.qwen3_next_rms_norm.Qwen3NextRMSNormOp"), + "cuda": "rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp", + "cuda-sm90": ("rl_engine.kernels.ops.cuda.norm.rmsnorm.Qwen3NextRMSNormCudaOp"), + }, + grad_input_names=("x", "weight"), + ), "qk_norm": OperatorSpec( name="qk_norm", op_class="reduction", diff --git a/tests/models/qwen3_next/__init__.py b/tests/models/qwen3_next/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/tests/models/qwen3_next/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/tests/models/qwen3_next/check_qwen3_next_norm_providers.py b/tests/models/qwen3_next/check_qwen3_next_norm_providers.py new file mode 100644 index 000000000..f4d25d2d8 --- /dev/null +++ b/tests/models/qwen3_next/check_qwen3_next_norm_providers.py @@ -0,0 +1,239 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""How far apart vLLM's eager gated-RMSNorm paths are, and ours from each. + +Named ``check_`` rather than ``test_`` on purpose, following +``tests/distributed/check_*.py``: this module imports real vLLM, and +The integration isolation test asserts that ``vllm`` is absent +from ``sys.modules`` -- an invariant any collected test importing vLLM would break +for the whole session. Run it explicitly: + + pytest tests/models/qwen3_next/check_qwen3_next_norm_providers.py -v + +RFC #428 section 2.1 defines L2 as "bitwise identical to vLLM rollout". For +Qwen3-Next's GDN gated norm that phrase is not well defined until a single +provider is named, because vLLM ships several and they do not agree bitwise with +each other. + +This module records two things so a vLLM upgrade cannot move them silently: + +1. **Provider facts** -- the GDN decode env defaults, and that ``RMSNormGated`` + has distinct ``forward_native`` and ``forward_cuda`` methods. It does NOT + determine which path vLLM dispatches at runtime: the checks below call each + method directly (eager), and vLLM's default compiled mode traces + ``forward_native`` into an inductor graph instead. +2. **The size of the gap** -- a seed sweep that asserts an upper bound on the + disagreement and on the mismatch rate between the eager paths. The bounds are + provider-gap bounds, not ``tolerance_contract.json`` thresholds, and they + deliberately do NOT assert equality. + +Measured on 2x B200 (sm_100, torch 2.13.0+cu130, vllm 0.30.0), bf16, +``head_v_dim=128``, 512 rows, 40 seeds: + +=============================== ================== ================ +comparison seeds not bitwise worst max|diff| +=============================== ================== ================ +ours vs ``forward_native`` 6 / 40 1.56e-2 +ours vs ``forward_cuda`` 18 / 40 3.91e-3 +``forward_native`` vs ``cuda`` 21 / 40 1.56e-2 +=============================== ================== ================ +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormGatedOp + +# vLLM is imported inside the checks, never at module scope, so that merely +# collecting this file does not pull it into sys.modules. + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") + +# Qwen3-Next-80B-A3B-Instruct: linear_value_head_dim / rms_norm_eps. +_HEAD_V_DIM = 128 +_EPS = 1e-6 +_ROWS = 512 +_SEEDS = 40 + +# The gated norm is constructed by vLLM's GDN block with exactly these settings; +# see vllm/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py. +_NORM_BEFORE_GATE = True +_GROUP_SIZE = None +_ACTIVATION = "silu" + + +@pytest.fixture(scope="module") +def vllm_config_ctx(): + """`RMSNormGated` is a CustomOp and refuses to build outside a config context.""" + pytest.importorskip("vllm", reason="vLLM is required to identify the provider") + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + yield + + +def _make_norm(): + from vllm.model_executor.layers.layernorm import RMSNormGated + + return RMSNormGated( + _HEAD_V_DIM, + eps=_EPS, + group_size=_GROUP_SIZE, + norm_before_gate=_NORM_BEFORE_GATE, + activation=_ACTIVATION, + ) + + +def _inputs(seed: int, dtype: torch.dtype, rows: int = _ROWS): + g = torch.Generator(device="cuda").manual_seed(seed) + x = torch.randn(rows, _HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + gate = torch.randn(rows, _HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + weight = torch.randn(_HEAD_V_DIM, device="cuda", dtype=dtype, generator=g) + return x, gate, weight + + +def _disagreement(a: torch.Tensor, b: torch.Tensor) -> tuple[float, float]: + """(worst absolute difference, fraction of elements that differ bitwise).""" + bits_a = a.float().view(torch.int32) + bits_b = b.float().view(torch.int32) + mismatch = int((bits_a != bits_b).sum()) + worst = (a.float() - b.float()).abs().max().item() + return worst, mismatch / a.numel() + + +# --------------------------------------------------------------------------- # +# 1. Provider facts -- recorded, not a runtime dispatch check +# --------------------------------------------------------------------------- # +def test_custom_op_has_distinct_native_and_cuda_paths(vllm_config_ctx): + """The two methods are distinct. Which one runs depends on the vLLM mode: eager + (``custom_ops="all"``) dispatches ``forward_cuda``; the default compiled mode + traces ``forward_native``. This test does not check that choice.""" + norm = _make_norm() + assert type(norm).forward_cuda is not type(norm).forward_native + + +def test_gdn_decode_provider_env_defaults_are_recorded(): + """Record the GDN decode env defaults. + + This records the defaults only; it does not assert which kernel runs. For + Qwen3-Next the ``VLLM_GDN_DECODE_KERNEL="cuda"`` default does not take effect: + vLLM builds its GDN layers with ``gqa_interleaved_layout=True`` and falls back + to the Triton decode kernel. A change of either default still fails here, as a + prompt to re-derive the decode path. + """ + pytest.importorskip("vllm", reason="vLLM is required to identify the provider") + import vllm.envs as envs + + observed = { + "VLLM_GDN_DECODE_KERNEL": envs.VLLM_GDN_DECODE_KERNEL.strip().lower(), + "VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE": bool( + envs.VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE + ), + } + assert observed == { + "VLLM_GDN_DECODE_KERNEL": "cuda", + "VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE": True, + }, ( + f"GDN decode provider defaults changed: {observed}. Re-derive which kernel " + "the rollout decode path takes before relying on any exactness claim." + ) + + +# --------------------------------------------------------------------------- # +# 2. The gap, bounded -- never asserted to be zero +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize( + "dtype, max_abs, max_mismatch_rate", + [ + # Provider-gap bounds, not contract thresholds. bf16 rounding absorbs most + # of the reduction-tree difference, so few elements move. + (torch.bfloat16, 2e-2, 0.05), + # In fp32 about 36% of elements differ, so the rate is left unbounded and + # only the magnitude is bounded (no per-element ULP bound is asserted). + (torch.float32, 1e-5, 1.0), + ], +) +def test_vllm_paths_disagree_only_within_bounds(vllm_config_ctx, dtype, max_abs, max_mismatch_rate): + """vLLM's own two paths differ; bound how much. + + A growing magnitude means the reduction trees have diverged semantically, + which would invalidate treating either as the reference. A high mismatch + *rate* at ULP magnitude does not -- it just means the tree shapes differ. + """ + norm = _make_norm().to("cuda", dtype) + worst_abs, worst_rate, differing = 0.0, 0.0, 0 + for seed in range(_SEEDS): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + native = norm.forward_native(x, gate) + cuda = norm.forward_cuda(x, gate) + abs_d, rate = _disagreement(native, cuda) + worst_abs, worst_rate = max(worst_abs, abs_d), max(worst_rate, rate) + differing += int(rate > 0.0) + + assert worst_abs <= max_abs, ( + f"vLLM forward_native vs forward_cuda worst |diff| {worst_abs:.3e} exceeds " + f"{max_abs:.3e} over {_SEEDS} seeds ({differing} seeds differ)" + ) + assert worst_rate <= max_mismatch_rate + + +@pytest.mark.parametrize("path", ["forward_native", "forward_cuda"]) +def test_ours_tracks_each_vllm_path_within_bounds(vllm_config_ctx, path): + """Our strict op follows vLLM's convention; bound the residual. + + Not an equality assertion. We reproduce the fp32 weight multiply and the + single trailing cast, but our reduction is the repo's fixed 32-wide chunked + sum rather than whatever tree the provider uses, so a few elements straddle + a rounding boundary. + """ + dtype = torch.bfloat16 + ours = Qwen3NextRMSNormGatedOp() + norm = _make_norm().to("cuda", dtype) + + worst_abs, worst_rate = 0.0, 0.0 + for seed in range(_SEEDS): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + reference = getattr(norm, path)(x, gate) + abs_d, rate = _disagreement(ours.forward(x, weight, gate), reference) + worst_abs, worst_rate = max(worst_abs, abs_d), max(worst_rate, rate) + + # Provider-gap bounds, not contract thresholds. + assert worst_abs <= 2e-2, f"worst |diff| vs {path} was {worst_abs:.3e}" + assert worst_rate <= 0.05, f"mismatch rate vs {path} was {worst_rate:.3%}" + + +# --------------------------------------------------------------------------- # +# 3. Batch invariance -- the property we DO claim +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("path", ["forward_native", "forward_cuda"]) +def test_vllm_gated_norm_is_batch_invariant(vllm_config_ctx, path): + """If this ever fails, vLLM stops being a coherent L2 target. + + Kept as a guard rather than a claim about our code: a provider whose output + depends on unrelated rows cannot anchor a bitwise contract. + """ + dtype = torch.bfloat16 + norm = _make_norm().to("cuda", dtype) + for seed in range(8): + x, gate, weight = _inputs(seed, dtype) + norm.weight.data = weight.clone() + full = getattr(norm, path)(x, gate) + for n in (1, 2, 8, 16, 32, 48, 64, 256): + sliced = getattr(norm, path)(x[:n], gate[:n]) + assert torch.equal(sliced, full[:n]), f"{path} seed={seed} n={n}" + + +def test_our_gated_op_is_batch_invariant(): + """The L1 claim for this operator, bitwise.""" + dtype = torch.bfloat16 + ours = Qwen3NextRMSNormGatedOp() + for seed in range(8): + x, gate, weight = _inputs(seed, dtype) + full = ours.forward(x, weight, gate) + for n in (1, 2, 8, 16, 32, 48, 64, 256): + assert torch.equal(ours.forward(x[:n], weight, gate[:n]), full[:n]) diff --git a/tests/models/qwen3_next/test_qwen3_next_norm.py b/tests/models/qwen3_next/test_qwen3_next_norm.py new file mode 100644 index 000000000..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_norm_reuse_check.py b/tests/models/qwen3_next/test_qwen3_next_norm_reuse_check.py new file mode 100644 index 000000000..d0140cec6 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_norm_reuse_check.py @@ -0,0 +1,125 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Keep providers with incompatible weight semantics out of norm evidence.""" + +import sys +import types + +import pytest +import torch + +from tools.validation.models import qwen3_next_norm_reuse_check as reuse + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("uses_effective_weight", [False, True]) +def test_megatron_zero_centered_admission(monkeypatch, dtype, uses_effective_weight): + """Reject the 52fbcbc forward behavior while admitting a corrected provider.""" + module = types.ModuleType("megatron.core.transformer.custom_layers.batch_invariant_kernels") + + class BatchInvariantRMSNormFn: + @staticmethod + def apply(x, weight, eps, zero_centered_gamma): + weight_eff = weight + 1.0 if zero_centered_gamma else weight + scale = weight_eff if uses_effective_weight else weight + normalized = x.float() * torch.rsqrt(x.float().square().mean(-1, keepdim=True) + eps) + return (normalized * scale.float()).to(x.dtype) + + module.BatchInvariantRMSNormFn = BatchInvariantRMSNormFn + monkeypatch.setitem(sys.modules, module.__name__, module) + monkeypatch.setattr(reuse, "DEV", "cpu") + monkeypatch.setattr(reuse, "DT", dtype) + factory = dict(reuse._c1_candidates(32))["megatron"] + monkeypatch.setattr(reuse, "_c1_candidates", lambda hidden: [("megatron", factory)]) + + candidate = reuse.candidates("qwen3_next_rms_norm")["megatron"] + if not uses_effective_weight: + assert "failed the zero-weight probe" in candidate["unavailable"] + assert "fn" not in candidate + return + + assert candidate["backward"] == "full" + x = torch.ones(2, 32, dtype=dtype, requires_grad=True) + weight = torch.zeros(32, dtype=dtype, requires_grad=True) + output = candidate["fn"](x, weight) + assert torch.all(output != 0) + output.sum().backward() + assert x.grad is not None + assert weight.grad is not None + + +@pytest.mark.parametrize("op", ["qwen3_next_rms_norm", "rms_norm_gated"]) +def test_required_candidate_failure_writes_no_report(monkeypatch, tmp_path, op): + """A missing primary extension must fail before creating report artifacts.""" + + def unavailable(): + raise RuntimeError("CUDA extension needs rebuilding") + + factory_name = "_c1_candidates" if op == "qwen3_next_rms_norm" else "_gated_candidates" + monkeypatch.setattr(reuse, factory_name, lambda hidden: [("rl_kernel", unavailable)]) + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(reuse, "plot", lambda *args: pytest.fail("must not plot a failed run")) + report = tmp_path / "evidence" / "report.json" + monkeypatch.setattr(sys, "argv", ["reuse", "--op", op, "--out", str(report)]) + + with pytest.raises(RuntimeError, match="Required rl_kernel candidate is unavailable"): + reuse.main() + + assert not report.parent.exists() + + +@pytest.mark.parametrize("op", ["qwen3_next_rms_norm", "rms_norm_gated"]) +def test_optional_candidate_failure_is_reported(monkeypatch, op): + """Optional dependencies may be absent when the primary candidate is usable.""" + + def primary(): + return "rl-kernel", lambda x, w: x, "full" + + def optional(): + raise ImportError("optional provider not installed") + + factory_name = "_c1_candidates" if op == "qwen3_next_rms_norm" else "_gated_candidates" + monkeypatch.setattr( + reuse, factory_name, lambda hidden: [("rl_kernel", primary), ("optional", optional)] + ) + found = reuse.candidates(op) + assert callable(found["rl_kernel"]["fn"]) + assert found["optional"] == {"unavailable": "ImportError: optional provider not installed"} + + +@pytest.mark.parametrize("op", ["qwen3_next_rms_norm", "rms_norm_gated"]) +@pytest.mark.parametrize("defect", [None, "forward", "gradient"]) +def test_full_batch_sub_batches_visit_previously_skipped_rows(monkeypatch, op, defect): + """Catch a singleton-only defect at an odd row missed by the old stride.""" + monkeypatch.setattr(reuse, "DEV", "cpu") + monkeypatch.setattr(reuse, "DT", torch.float32) + monkeypatch.setitem(reuse.SHAPES, op, (1, (8,), 1025, (1, 7, 64), 1025)) + + def inputs(op, seed, n): + tensors = {"x": torch.arange(n, dtype=torch.float32).reshape(n, 1), "w": torch.ones(1)} + if op == "rms_norm_gated": + tensors["z"] = torch.ones(n, 1) + return tensors + + monkeypatch.setattr(reuse, "_inputs", inputs) + + def provider(x, w, z=None): + output = x * w + if z is not None: + output = output * z + if len(x) == 1: + affected = (x == 1001).float() + if defect == "forward": + output = output + affected + elif defect == "gradient": + output = output + (x - x.detach()) * affected + return output + + result = reuse.batch_invariance(op, {"fn": provider, "backward": "x-only"}, quick=False) + check = result["full_vs_sub_batches"] + assert check["coverage"] == "every_row_per_sub_batch_size" + assert check["sub_batches"] == 2 * (1025 + 147 + 17) + assert check["fwd_differ"] == (2 if defect == "forward" else 0) + assert check["rowgrad_differ"] == (2 if defect == "gradient" else 0) + assert result["batch_invariant"] is (defect is None) diff --git a/tests/models/qwen3_next/test_qwen3_next_workload.py b/tests/models/qwen3_next/test_qwen3_next_workload.py new file mode 100644 index 000000000..9710f0151 --- /dev/null +++ b/tests/models/qwen3_next/test_qwen3_next_workload.py @@ -0,0 +1,66 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The norm workload cannot stand in for Dense or full-checkpoint evidence.""" + +import copy +from pathlib import Path + +import pytest + +from rl_engine.config.workload import WorkloadError, load_manifest, manifest_identity_hash +from rl_engine.validation.models.qwen3_next_workload import validate_norm_manifest +from rl_engine.validation.operators.gradient_adapters import get_adapter, resolve_profile_candidate + +MANIFEST = ( + Path(__file__).resolve().parents[3] + / "rl_engine/validation/models/qwen3_next_norm_manifest.json" +) + + +def test_norm_manifest_pins_real_architecture_and_separate_scope(): + manifest = load_manifest(MANIFEST) + assert manifest.model_identity["config_fingerprint"]["hidden_size"] == 2048 + assert manifest.raw["full_model_evidence"] is False + for name, gradients in ( + ("qwen3_next_rms_norm", ("dx", "dweight")), + ("rms_norm_gated", ("dx", "dweight", "dgate")), + ): + adapter = get_adapter(name) + assert tuple(t.name for t in adapter.tensors) == gradients + assert resolve_profile_candidate(adapter, "cuda_bf16", manifest)["status"] == "declared" + assert ( + resolve_profile_candidate(adapter, "triton_cuda_bf16", manifest)["status"] + == "missing_required" + ) + assert ( + resolve_profile_candidate(adapter, "cuda_bf16", load_manifest())["status"] + == "absent_not_required" + ) + + +@pytest.mark.parametrize("fault", ["architecture", "full_model", "revision", "shape", "binding"]) +def test_norm_manifest_rejects_false_evidence_even_with_regenerated_hash(fault): + raw = copy.deepcopy(load_manifest(MANIFEST).raw) + if fault == "architecture": + raw["model_identity"]["config_fingerprint"]["hidden_size"] = 4096 + elif fault == "full_model": + raw["full_model_evidence"] = True + elif fault == "revision": + raw["model_identity"]["revision"] = "main" + elif fault == "shape": + raw["representative_cases"][0]["hidden"] = 64 + else: + raw["representative_cases"][0]["fixture_id"] = "wrong_fixture" + raw["fixture_identity_sha256"] = manifest_identity_hash(raw) + with pytest.raises(WorkloadError): + validate_norm_manifest(raw) + + +def test_norm_dimension_gate_rejects_shrunk_workload(): + from rl_engine.validation.models.qwen3_next_workload import validate_norm_dimensions + + raw = load_manifest(MANIFEST).raw + with pytest.raises(WorkloadError, match="hidden 2048"): + validate_norm_dimensions(raw, "qwen3_next_rms_norm", 64, 128) + validate_norm_dimensions(raw, "rms_norm_gated", 2048, 128) diff --git a/tests/ops/norm/test_rms_norm.py b/tests/ops/norm/test_rms_norm.py index 90d7bab23..352612805 100644 --- a/tests/ops/norm/test_rms_norm.py +++ b/tests/ops/norm/test_rms_norm.py @@ -17,10 +17,11 @@ try: from rl_engine.backends.extension import _C, _EXT_AVAILABLE - # The same two symbols RMSNormCudaOp.__init__ requires; a build that has them - # dispatches to the CUDA op, so the dispatch test must agree with that guard. - _HAS_CUDA_RMSNORM = _EXT_AVAILABLE and all( - hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx") + # Keep dispatch expectations aligned with the constructor's capability guard. + _HAS_CUDA_RMSNORM = ( + _EXT_AVAILABLE + and getattr(_C, "rmsnorm_api_version", None) == 2 + and all(hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx")) ) except ImportError: # pragma: no cover - import can fail when the extension is not built. _HAS_CUDA_RMSNORM = False @@ -251,7 +252,7 @@ def test_backward_batch_invariance_slice(): # 9b. The CUDA backend must report itself unavailable by failing construction, # which is the seam the registry uses to fall back (see _get_or_create_backend). def test_cuda_op_construction_fails_without_extension(monkeypatch): - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) monkeypatch.setattr(cuda_rmsnorm, "_C", None) @@ -260,10 +261,10 @@ def test_cuda_op_construction_fails_without_extension(monkeypatch): def test_cuda_op_construction_fails_when_symbols_missing(monkeypatch): - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm class _WithoutRMSNorm: # a built extension that lacks the rmsnorm symbols - pass + rmsnorm_api_version = 2 monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr(cuda_rmsnorm, "_C", _WithoutRMSNorm()) @@ -273,8 +274,8 @@ class _WithoutRMSNorm: # a built extension that lacks the rmsnorm symbols def test_registry_falls_back_to_native_without_extension(monkeypatch): """A CUDA-first priority list must still resolve on a build without _C.""" - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) monkeypatch.setattr(cuda_rmsnorm, "_C", None) @@ -286,30 +287,66 @@ def test_registry_falls_back_to_native_without_extension(monkeypatch): def test_registry_falls_back_when_required_symbol_is_missing(monkeypatch, missing): from types import SimpleNamespace - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry symbols = {name: object() for name in ("rmsnorm_forward", "rmsnorm_backward_dx")} del symbols[missing] + symbols["rmsnorm_api_version"] = 2 monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr(cuda_rmsnorm, "_C", SimpleNamespace(**symbols)) assert isinstance(KernelRegistry().get_op("rms_norm", device="cuda"), NativeRMSNormOp) -def test_registry_cuda_requires_only_used_symbols_and_cpu_stays_native(monkeypatch): +@pytest.mark.parametrize("api_version", [None, 1, 3]) +def test_cuda_rejects_incompatible_rmsnorm_api_before_dispatch(monkeypatch, api_version): + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry + + class _LegacyRMSNorm: + def rmsnorm_forward(self, x, weight, eps): + pytest.fail("an incompatible extension must not be invoked") + + def rmsnorm_backward_dx(self, dy, x, weight, rstd): + pytest.fail("an incompatible extension must not be invoked") + + extension = _LegacyRMSNorm() + if api_version is not None: + extension.rmsnorm_api_version = api_version + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) + monkeypatch.setattr(cuda_rmsnorm, "_C", extension) + + for op_type in (cuda_rmsnorm.RMSNormCudaOp, cuda_rmsnorm.Qwen3NextRMSNormCudaOp): + with pytest.raises(RuntimeError, match="RMSNorm API version 2.*Rebuild"): + op_type() + + registry_op = KernelRegistry().get_op("rms_norm", device="cuda") + assert isinstance(registry_op, NativeRMSNormOp) + x, weight = torch.ones(2, 8), torch.ones(8) + assert torch.isfinite(registry_op(x, weight)).all() + with pytest.raises(RuntimeError, match="RMSNorm API version 2.*Rebuild"): + rmsnorm_cuda(x, weight, weight_offset=1.0) + + +def test_registry_cuda_requires_current_api_and_only_used_symbols(monkeypatch): from types import SimpleNamespace - from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm - from rl_engine.kernels.registry import KernelRegistry + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.runtime.registry import KernelRegistry monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) monkeypatch.setattr( cuda_rmsnorm, "_C", - SimpleNamespace(rmsnorm_forward=object(), rmsnorm_backward_dx=object()), + SimpleNamespace( + rmsnorm_api_version=2, + rmsnorm_forward=object(), + rmsnorm_backward_dx=object(), + ), ) registry = KernelRegistry() assert isinstance(registry.get_op("rms_norm", device="cuda"), RMSNormCudaOp) + cuda_rmsnorm.Qwen3NextRMSNormCudaOp() assert isinstance(registry.get_op("rms_norm", device="cpu"), NativeRMSNormOp) diff --git a/tests/runtime/test_dispatch.py b/tests/runtime/test_dispatch.py index 3c47bbbcf..6517d3309 100644 --- a/tests/runtime/test_dispatch.py +++ b/tests/runtime/test_dispatch.py @@ -4,6 +4,7 @@ import sys from types import ModuleType +import pytest import torch import rl_engine.platforms.device as device_module @@ -258,3 +259,67 @@ def test_executor_flow(): print("\n All infrastructure tests passed!") except Exception as e: print(f"\n Test failed with error: {e}") + + +_QWEN3_NEXT_NORMS = [ + pytest.param( + "rms_norm_gated", + OpBackend.CUDA_RMS_NORM_GATED, + OpBackend.PYTORCH_NATIVE_RMS_NORM_GATED, + "Qwen3NextRMSNormGatedCudaOp", + "Qwen3NextRMSNormGatedOp", + id="rms_norm_gated", + ), + pytest.param( + "qwen3_next_rms_norm", + OpBackend.CUDA_QWEN3_NEXT_RMS_NORM, + OpBackend.PYTORCH_NATIVE_QWEN3_NEXT_RMS_NORM, + "Qwen3NextRMSNormCudaOp", + "Qwen3NextRMSNormOp", + id="qwen3_next_rms_norm", + ), +] + + +@pytest.mark.parametrize("op_type, cuda_backend, ref_backend, cuda_cls, ref_cls", _QWEN3_NEXT_NORMS) +def test_qwen3_next_norm_priority_is_cuda_first_with_pytorch_fallback( + op_type, cuda_backend, ref_backend, cuda_cls, ref_cls +): + """Both Qwen3-Next norms have a CUDA kernel; every other platform falls back. + + No Triton, ROCm or Ascend kernel exists yet, so those platforms must resolve + to the PyTorch reference rather than to nothing -- an operator missing from a + priority map falls through to ``OpBackend.PYTORCH_NATIVE``, which is the + logprob op, not a norm. + """ + registry = KernelRegistry() + + assert registry._priority_map["cuda"][op_type] == [cuda_backend, ref_backend] + for platform in ("rocm", "musa", "cpu", "npu"): + assert registry._priority_map[platform][op_type] == [ref_backend], platform + + +@pytest.mark.skipif(torch.version.hip is not None, reason="a cuda device maps to rocm on HIP") +@pytest.mark.parametrize("op_type, cuda_backend, ref_backend, cuda_cls, ref_cls", _QWEN3_NEXT_NORMS) +def test_qwen3_next_norm_cuda_backend_absence_falls_through_to_reference( + monkeypatch, op_type, cuda_backend, ref_backend, cuda_cls, ref_cls +): + """Without the compiled symbols the CUDA backend must be skipped, not returned. + + The zero-centred op inherits its check from ``RMSNormCudaOp.__init__``. The + lookup names a CUDA device so the CUDA-first list is walked even on a + CPU-only host; resolving for the host's own platform would reach the + reference through the CPU list without ever trying the CUDA backend. + """ + from rl_engine.backends.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.reference.norm import qwen3_next_rms_norm as reference + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) + monkeypatch.setattr(cuda_rmsnorm, "_C", None) + with pytest.raises(RuntimeError, match="requires the compiled rl_engine._C extension"): + getattr(cuda_rmsnorm, cuda_cls)() + + registry = KernelRegistry() + resolved = registry.get_op(op_type, device="cuda") + assert type(resolved) is getattr(reference, ref_cls) + assert cuda_backend.name in registry._failed_backends diff --git a/tests/validation/operators/test_ws1_ascend_closeout.py b/tests/validation/operators/test_ws1_ascend_closeout.py index 66c8ae4a6..b9b0fe331 100644 --- a/tests/validation/operators/test_ws1_ascend_closeout.py +++ b/tests/validation/operators/test_ws1_ascend_closeout.py @@ -187,9 +187,15 @@ def test_c2_ascend_cases_pin_real_ascend_kernels_and_sources(): def test_c3_c4_every_required_adapter_resolves_an_ascend_candidate(): manifest = load_manifest() + model_id = manifest.model_identity["model_id"] for name, adapter in GRADIENT_ADAPTERS.items(): if adapter.requirement not in ("required",): continue + # An adapter scoped to another model (the CUDA-only Qwen3-Next norms) + # resolves to `absent_not_required` for this manifest by design; this + # test covers the chain of the model whose manifest it loads. + if adapter.model_id not in (None, model_id): + continue resolved = resolve_profile_candidate(adapter, PROFILE, manifest) assert resolved["status"] == "declared", name assert resolved["expected_backend_id"] == "ascend", name @@ -197,6 +203,22 @@ def test_c3_c4_every_required_adapter_resolves_an_ascend_candidate(): assert candidate_family(str(resolved["expected_backend_id"])) == "ascend" +def test_model_scoped_adapters_are_absent_not_required_for_other_models(): + """The scoping the test above relies on must actually hold. + + Without this, skipping scoped adapters could hide one that silently resolves + to a real candidate for the wrong model. + """ + manifest = load_manifest() + model_id = manifest.model_identity["model_id"] + scoped = [a for a in GRADIENT_ADAPTERS.values() if a.model_id not in (None, model_id)] + assert scoped, "expected at least the Qwen3-Next norm adapters to be model-scoped" + for adapter in scoped: + resolved = resolve_profile_candidate(adapter, PROFILE, manifest) + assert resolved["status"] == "absent_not_required", adapter.op_name + assert resolved["candidate_path"] is None, adapter.op_name + + def test_c4_adapter_status_matrix_has_no_red_ascend_rows(): rows = [r for r in gradient_adapter_status_matrix() if r.backend_profile == PROFILE] assert rows diff --git a/tests/validation/operators/test_ws1_gtest_gpu.py b/tests/validation/operators/test_ws1_gtest_gpu.py index bcaebc2f5..60ac5102c 100644 --- a/tests/validation/operators/test_ws1_gtest_gpu.py +++ b/tests/validation/operators/test_ws1_gtest_gpu.py @@ -36,6 +36,8 @@ def test_all_ws1_single_ops_are_registered(): names = set(operator_names()) assert { "rms_norm", + "rms_norm_gated", + "qwen3_next_rms_norm", "qk_norm", "det_gemm", "attention", diff --git a/tools/validation/models/plot_qwen3_next_norm_evidence.py b/tools/validation/models/plot_qwen3_next_norm_evidence.py new file mode 100644 index 000000000..ede305a6f --- /dev/null +++ b/tools/validation/models/plot_qwen3_next_norm_evidence.py @@ -0,0 +1,122 @@ +#!/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 -m pip install matplotlib # optional plotting dependency +python tools/validation/models/plot_qwen3_next_norm_evidence.py report.json + +Writes figure[-].png beside the report. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 + +COLORS = ["#2a6fdb", "#7a7a7a", "#e07b39", "#3aa676", "#9b59b6", "#c0392b"] + + +def _short(name: str) -> str: + return name.replace(" (forward only)", "*").replace(" (cast-first)", "\n(cast-first)") + + +def plot_op(op: str, data: dict, title: str, out: Path) -> None: + names = list(data["row_invariance"]) + color = {n: COLORS[i % len(COLORS)] for i, n in enumerate(names)} + fig, axes = plt.subplots(2, 2, figsize=(14, 10), layout="constrained") + fig.suptitle(title, fontsize=13) + + for ax, key, label in ( + (axes[0, 0], "forward_us", "forward latency"), + (axes[0, 1], "backward_us", "backward latency (training-capable only)"), + ): + for name in names: + rows = sorted(int(r) for r in data["latency"][name]) + ys = [data["latency"][name][str(r)].get(key) for r in rows] + if all(y is None for y in ys): + continue + ax.plot(rows, ys, "o-", color=color[name], label=_short(name)) + ax.set_xscale("log", base=2) + ax.set_yscale("log") + ax.set_xlabel("rows") + ax.set_ylabel("µs (median)") + ax.set_title(label) + ax.grid(True, which="both", alpha=0.3) + ax.legend(fontsize=8) + + ax = axes[1, 0] + acc = data["accuracy"][max(data["accuracy"], key=int)] + metrics = [ + ("forward_max_abs", "forward\nmax |err|"), + ("dx_max_abs_over_absmax", "dx\nmax |err| / max"), + ("dweight_max_abs_over_absmax", "dweight\nmax |err| / max"), + ("dgate_max_abs_over_absmax", "dgate\nmax |err| / max"), + ] + metrics = [m for m in metrics if any(acc[n].get(m[0]) is not None for n in names)] + width = 0.8 / len(names) + for i, name in enumerate(names): + vals = [acc[name].get(m) for m, _ in metrics] + xs = [j + (i - (len(names) - 1) / 2) * width for j in range(len(metrics))] + ax.bar( + [x for x, v in zip(xs, vals) if v is not None], + [v for v in vals if v is not None], + width, + color=color[name], + label=_short(name), + ) + ax.set_xticks(range(len(metrics)), [label for _, label in metrics], fontsize=9) + ax.set_yscale("log") + ax.set_title(f"error vs FP64 golden ({max(data['accuracy'], key=int)} rows, BF16)") + ax.grid(True, axis="y", alpha=0.3) + ax.legend(fontsize=8) + + ax = axes[1, 1] + bi = data["row_invariance"] + checked = next(iter(bi.values()))["rows_checked"] + ys = range(len(names)) + fwd = [bi[n]["forward_rows_differing"] for n in names] + dx = [ + bi[n]["dx_rows_differing"] if bi[n]["dx_rows_differing"] is not None else 0 for n in names + ] + ax.barh([y - 0.2 for y in ys], fwd, 0.4, color="#2a6fdb", label="forward") + ax.barh([y + 0.2 for y in ys], dx, 0.4, color="#e07b39", label="dx") + for y, n, f, d in zip(ys, names, fwd, dx): + no_bwd = bi[n]["dx_rows_differing"] is None + ax.text(max(f, d) + 0.2, y, f"{f} / {'n/a' if no_bwd else d}", va="center", fontsize=8) + ax.set_yticks(list(ys), [_short(n) for n in names], fontsize=8) + ax.set_xlim(0, max(3, max(fwd + dx) * 1.4)) + ax.invert_yaxis() + ax.set_xlabel(f"rows differing (of {checked}; row alone vs inside a batch, bitwise)") + ax.set_title("row invariance (0 = batch-invariant)") + ax.legend(fontsize=8) + ax.grid(True, axis="x", alpha=0.3) + + fig.savefig(out, dpi=130) + print(f"wrote {out}") + + +def main() -> None: + path = Path(sys.argv[1]) + report = json.loads(path.read_text()) + env = report["environment"] + ops = report["ops"] + for op, data in ops.items(): + name = "figure.png" if len(ops) == 1 else f"figure-{op}.png" + hidden = data.get("hidden", report.get("hidden")) + title = ( + f"Qwen3-Next {op.replace('_', ' ')} — {env['gpu']}, hidden {hidden}, " + f"BF16, commit {report['git_commit'][:7]} (* forward only)" + ) + plot_op(op, data, title, path.parent / name) + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/qwen3_next_norm_evidence.py b/tools/validation/models/qwen3_next_norm_evidence.py new file mode 100644 index 000000000..aaeba82c3 --- /dev/null +++ b/tools/validation/models/qwen3_next_norm_evidence.py @@ -0,0 +1,348 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Accuracy, row invariance and latency of the Qwen3-Next norms vs existing implementations. + +Writes one JSON report (RFC #428, reuse rule of #420): every candidate is compared with +an FP64 golden, checked for row invariance (a row computed alone vs inside a batch, +bitwise), and timed. Optional providers (transformers, vLLM, FlashInfer) are skipped +when they are not installed; the report says which ran. + + python tools/validation/models/qwen3_next_norm_evidence.py --out report.json + python tools/validation/models/plot_qwen3_next_norm_evidence.py report.json +""" + +from __future__ import annotations + +import argparse +import json +import platform +import statistics +import subprocess +import sys +from pathlib import Path +from typing import Any, Callable + +import torch +import torch.nn.functional as F + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.backends.cuda.norm.rmsnorm import ( # noqa: E402 + Qwen3NextRMSNormCudaOp, + Qwen3NextRMSNormGatedCudaOp, +) +from rl_engine.reference.norm.qwen3_next_rms_norm import ( # noqa: E402 + Qwen3NextRMSNormGatedOp, + Qwen3NextRMSNormOp, +) + +EPS = 1e-6 + + +# --------------------------------------------------------------------------- # +# Ops. Each candidate is fn(*row_inputs, weight) -> y, plus whether it has a +# backward. Row inputs are sliced together for the row-invariance check. +# --------------------------------------------------------------------------- # + + +def zero_centred_candidates(hidden: int) -> dict[str, dict[str, Any]]: + cands: dict[str, dict[str, Any]] = { + "rl-kernel CUDA": {"fn": Qwen3NextRMSNormCudaOp(), "backward": True}, + "rl-kernel PyTorch reference": {"fn": Qwen3NextRMSNormOp(), "backward": True}, + } + try: + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNorm + + module = Qwen3NextRMSNorm(hidden, eps=EPS).cuda() + + def hf(x, w, module=module): + return torch.func.functional_call(module, {"weight": w}, (x,)) + + cands["transformers Qwen3NextRMSNorm"] = {"fn": hf, "backward": True} + except ImportError: + pass + try: + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + from vllm.model_executor.layers.layernorm import GemmaRMSNorm + + gemma = GemmaRMSNorm(hidden, eps=EPS).cuda() + + def vllm_fwd(x, w, gemma=gemma): + gemma.weight.data = w.detach() + return gemma.forward_cuda(x) + + cands["vLLM GemmaRMSNorm (forward only)"] = {"fn": vllm_fwd, "backward": False} + except ImportError: + pass + try: + import flashinfer + + cands["FlashInfer gemma_rmsnorm (forward only)"] = { + "fn": lambda x, w: flashinfer.norm.gemma_rmsnorm(x, w, EPS), + "backward": False, + } + except ImportError: + pass + return cands + + +def zero_centred_golden(x, w): + x64 = x.double() + return x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + EPS) * (1.0 + w.double()) + + +def gated_candidates(hidden: int) -> dict[str, dict[str, Any]]: + cuda_op, ref_op = Qwen3NextRMSNormGatedCudaOp(), Qwen3NextRMSNormGatedOp() + cands: dict[str, dict[str, Any]] = { + "rl-kernel CUDA": {"fn": lambda x, g, w: cuda_op(x, w, g, eps=EPS), "backward": True}, + "rl-kernel PyTorch reference": { + "fn": lambda x, g, w: ref_op(x, w, g, eps=EPS), + "backward": True, + }, + } + try: + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNormGated + + module = Qwen3NextRMSNormGated(hidden, eps=EPS).cuda() + + def hf(x, g, w, module=module): + return torch.func.functional_call(module, {"weight": w}, (x, g)) + + cands["transformers Qwen3NextRMSNormGated (cast-first)"] = {"fn": hf, "backward": True} + except ImportError: + pass + try: + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + from vllm.model_executor.layers.layernorm import RMSNormGated + + gated = RMSNormGated(hidden, eps=EPS, norm_before_gate=True).cuda() + + def vllm_fwd(x, g, w, gated=gated): + gated.weight.data = w.detach() + return gated.forward_cuda(x, g) + + cands["vLLM RMSNormGated (forward only)"] = {"fn": vllm_fwd, "backward": False} + except ImportError: + pass + return cands + + +def gated_golden(x, g, w): + """vLLM's convention (the one #468 implements) in FP64: x * rstd * w * silu(gate).""" + + x64 = x.double() + return (x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + EPS) * w.double()) * F.silu( + g.double() + ) + + +OPS: dict[str, dict[str, Any]] = { + "zero_centred_rmsnorm": { + "hidden": 2048, # decoder and final norms + "row_inputs": 1, + "weight": lambda h, gen: torch.randn(h, device="cuda", generator=gen) * 0.1, + "candidates": zero_centred_candidates, + "golden": zero_centred_golden, + "timed_rows": (1024, 4096, 16384, 65536), + }, + "gated_rmsnorm": { + "hidden": 128, # GDN value head dim; rows are tokens x heads + "row_inputs": 2, + "weight": lambda h, gen: 1.0 + torch.randn(h, device="cuda", generator=gen) * 0.1, + "candidates": gated_candidates, + "golden": gated_golden, + "timed_rows": (4096, 16384, 65536, 262144), + }, +} + + +# --------------------------------------------------------------------------- # +# Measurements +# --------------------------------------------------------------------------- # + + +def _inputs(spec, rows: int, seed: int, dtype=torch.bfloat16): + g = torch.Generator(device="cuda").manual_seed(seed) + h = spec["hidden"] + row_inputs = [(torch.randn(rows, h, device="cuda", generator=g) * 2).to(dtype)] + for _ in range(spec["row_inputs"] - 1): + row_inputs.append(torch.randn(rows, h, device="cuda", generator=g).to(dtype)) + w = spec["weight"](h, g).to(dtype) + up = torch.randn(rows, h, device="cuda", generator=g).to(dtype) + return row_inputs, w, up + + +def _grads(fn, row_inputs, w, up, dtype=None): + def leaf(t): + return (t if dtype is None else t.to(dtype)).detach().clone().requires_grad_(True) + + rl, wl = [leaf(t) for t in row_inputs], leaf(w) + out = fn(*rl, wl) + out.backward(up if dtype is None else up.to(dtype)) + return out.detach(), [t.grad for t in rl], wl.grad + + +def accuracy(spec, cands, rows: int, seed: int) -> dict[str, Any]: + row_inputs, w, up = _inputs(spec, rows, seed) + ref_out, ref_drows, ref_dw = _grads(spec["golden"], row_inputs, w, up, torch.float64) + result = {} + for name, c in cands.items(): + entry: dict[str, Any] = {} + if c["backward"]: + out, drows, dw = _grads(c["fn"], row_inputs, w, up) + pairs = [("dx", drows[0], ref_drows[0]), ("dweight", dw, ref_dw)] + if len(drows) > 1: + pairs.append(("dgate", drows[1], ref_drows[1])) + for key, got, ref in pairs: + err = (got.double() - ref).abs().max().item() + entry[f"{key}_max_abs_over_absmax"] = err / ref.abs().max().item() + else: + with torch.no_grad(): + out = c["fn"](*row_inputs, w) + err = (out.double() - ref_out).abs() + entry["forward_max_abs"] = err.max().item() + entry["forward_correctly_rounded_fraction"] = ( + (out == ref_out.to(out.dtype)).float().mean().item() + ) + result[name] = entry + return result + + +def row_invariance(spec, cands, seeds=(3, 4, 5), rows: int = 4096, step: int = 16): + """256 rows computed alone vs the same rows inside a batch, bitwise.""" + + result = {} + for name, c in cands.items(): + fwd_bad = dx_bad = checked = 0 + for seed in seeds: + row_inputs, w, up = _inputs(spec, rows, seed) + if c["backward"]: + full_out, full_drows, _ = _grads(c["fn"], row_inputs, w, up) + else: + with torch.no_grad(): + full_out = c["fn"](*row_inputs, w) + for i in range(0, rows, step): + part = [t[i : i + 1] for t in row_inputs] + if c["backward"]: + out, drows, _ = _grads(c["fn"], part, w, up[i : i + 1]) + dx_bad += not all(torch.equal(d[0], fd[i]) for d, fd in zip(drows, full_drows)) + else: + with torch.no_grad(): + out = c["fn"](*part, w) + fwd_bad += not torch.equal(out[0], full_out[i]) + checked += 1 + result[name] = { + "rows_checked": checked, + "forward_rows_differing": fwd_bad, + "dx_rows_differing": dx_bad if c["backward"] else None, + "batch_rows": rows, + "seeds": list(seeds), + } + return result + + +def _time_us(fn: Callable[[], Any], warmup: int = 10, iters: int = 50) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(iters): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1e3) + return statistics.median(samples) + + +def latency(spec, cands) -> dict[str, Any]: + result: dict[str, Any] = {name: {} for name in cands} + for rows in spec["timed_rows"]: + row_inputs, w, up = _inputs(spec, rows, seed=11) + for name, c in cands.items(): + row: dict[str, float] = {} + with torch.no_grad(): + row["forward_us"] = _time_us(lambda f=c["fn"]: f(*row_inputs, w)) + if c["backward"]: + leaves = [t.detach().clone().requires_grad_(True) for t in [*row_inputs, w]] + out = c["fn"](*leaves) + row["backward_us"] = _time_us( + lambda o=out, lv=leaves: torch.autograd.grad(o, lv, up, retain_graph=True) + ) + result[name][str(rows)] = row + return result + + +# --------------------------------------------------------------------------- # + + +def _git(*args: str) -> str: + try: + return subprocess.check_output(["git", *args], cwd=REPO_ROOT, text=True).strip() + except (OSError, subprocess.CalledProcessError): + return "" + + +def environment() -> dict[str, Any]: + env = { + "gpu": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "python": platform.python_version(), + } + for mod in ("transformers", "vllm", "flashinfer"): + try: + env[mod] = __import__(mod).__version__ + except ImportError: + env[mod] = None + return env + + +def run_op(name: str, spec) -> dict[str, Any]: + cands = spec["candidates"](spec["hidden"]) + print(f"[{name}] candidates: {', '.join(cands)}", flush=True) + report = { + "hidden": spec["hidden"], + "accuracy": {str(r): accuracy(spec, cands, r, seed=r) for r in (257, 4096)}, + "row_invariance": row_invariance(spec, cands), + "latency": latency(spec, cands), + } + for cand, entry in report["row_invariance"].items(): + print(f" BI {cand}: {entry}", flush=True) + return report + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--ops", default=",".join(OPS), help="comma list of ops") + args = parser.parse_args() + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + torch.backends.cuda.matmul.allow_tf32 = False + report = { + "kind": "qwen3_next_norm_evidence", + "rfc": "RL-Align/RL-Kernel#428", + "git_commit": _git("rev-parse", "HEAD") or "unknown", + "git_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": environment(), + "eps": EPS, + "dtype": "bfloat16", + "ops": {name: run_op(name, OPS[name]) for name in args.ops.split(",")}, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + print(f"wrote {args.out} (commit {report['git_commit'][:7]}, dirty={report['git_dirty']})") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/qwen3_next_norm_reuse_check.py b/tools/validation/models/qwen3_next_norm_reuse_check.py new file mode 100644 index 000000000..06b24f9d5 --- /dev/null +++ b/tools/validation/models/qwen3_next_norm_reuse_check.py @@ -0,0 +1,714 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Compare a Qwen3-Next norm with existing implementations: batch invariance, accuracy, gates. + +Checks, for the rl-kernel CUDA op and every existing implementation that imports here +(transformers, vLLM, FlashInfer, Liger, FLA, Transformer Engine, Megatron-LM from a source +tree): + +* ``bi``: bitwise batch invariance of the forward and the row gradients, and repeatability of + ``dweight``: every row computed alone vs inside full batches of three sizes (seeds 3-5); + the full workload-size batch vs sub-batches covering every row; eight probe rows at the + front, middle and back of batches of every size 1..9 and 2^k - 1, 2^k, 2^k + 1. Sparse probes + miss row-specific and large-batch dependence, which is why every row is checked. +* ``perf``: accuracy against FP64 (forward: fraction equal to the correctly rounded FP64 result; + gradients: max|err| / max|ref|) and median CUDA-event latency at the workload size. +* ``gates``: this repository's own C3/C4 gates + (``tools/validation/operators/check_forward_invariance.py`` and + ``tools/validation/operators/check_gradient_invariance.py``) with the norm manifest and the + arguments ``ci/scripts/run_ws1_gtest.sh`` uses), unchanged, with the CUDA candidate replaced by a + subclass of the rl-kernel op whose forward and backward call the other library. The subclass + keeps the op's ``parameter_vjp_contributions_fp32``, so the singleton-aggregate check + compares that library's ``dweight`` with the same FP32 row contributions. + + python tools/validation/models/qwen3_next_norm_reuse_check.py --op qwen3_next_rms_norm \\ + --out docs/usage/evidence/qwen3-next-norm-reuse-b200/qwen3_next_rms_norm.json \\ + [--checks bi,perf,gates] [--megatron-src /path/to/Megatron-LM] [--quick] + +Run on a clean tree on an otherwise idle GPU; the report records the commit, the tree state and +every library version. A figure is written next to the report when matplotlib is available. +""" + +from __future__ import annotations + +import argparse +import importlib.metadata as md +import json +import os +import statistics +import subprocess +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +import torch # noqa: E402 +import torch.nn.functional as F # noqa: E402 + +DEV, EPS, DT = "cuda", 1e-6, torch.bfloat16 +SHAPES = { + # op: (hidden, all-rows batch sizes, full batch, covering sub-batch sizes, perf rows) + "qwen3_next_rms_norm": (2048, (64, 257, 2048), 65536, (1, 7, 2048, 4097, 32768), 65536), + "rms_norm_gated": (128, (64, 257, 4097), 262144, (1, 7, 4097, 65536, 131072), 262144), +} +LIBS = ("torch", "triton", "transformers", "vllm", "flashinfer-python", "liger-kernel", "fla-core") +LIBS += ("transformer_engine", "megatron-core") +MANIFEST = "rl_engine/validation/models/qwen3_next_norm_manifest.json" + + +def _version(name: str) -> str | None: + try: + return md.version(name) + except md.PackageNotFoundError: + return None + + +# --------------------------------------------------------------------------- # +# Candidates: fn(x, w) for the zero-centred norm, fn(x, z, w) for the gated norm +# --------------------------------------------------------------------------- # + + +def _vllm_module(build): + from vllm.config import VllmConfig, set_current_vllm_config + + with set_current_vllm_config(VllmConfig()): + return build() + + +def _c1_candidates(hidden: int) -> list[tuple[str, Any]]: + def ours(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormCudaOp + + o = Qwen3NextRMSNormCudaOp() + return "rl-kernel Qwen3NextRMSNormCudaOp", lambda x, w: o(x, w, eps=EPS), "full" + + def reference(): + from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormOp + + o = Qwen3NextRMSNormOp() + return "rl-kernel PyTorch reference", lambda x, w: o(x, w, eps=EPS), "full" + + def transformers(): + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNorm + + m = Qwen3NextRMSNorm(hidden, eps=EPS).to(DEV, DT) + + def fn(x, w): + return torch.func.functional_call(m, {"weight": w}, (x,)) + + return f"transformers {_version('transformers')} Qwen3NextRMSNorm", fn, "full" + + def liger(): + from liger_kernel.ops.rms_norm import LigerRMSNormFunction + + def fn(x, w): + return LigerRMSNormFunction.apply(x, w, EPS, 1.0, "gemma", False) + + return f"Liger {_version('liger-kernel')} RMSNorm, offset 1, gemma", fn, "full" + + def fla(): + from fla.modules.layernorm import rms_norm + + def fn(x, w): + return rms_norm(x, (1.0 + w.float()).to(x.dtype), None, eps=EPS) + + return f"FLA {_version('fla-core')} rms_norm, weight passed as 1 + w", fn, "full" + + def te(): + import transformer_engine.pytorch as te_ + + m = te_.RMSNorm(hidden, eps=EPS, zero_centered_gamma=True, params_dtype=DT, device=DEV) + + def fn(x, w): + # TE's backward reads its own parameter, so copy w in (functional_call gives + # a wrong dx); dweight is then TE's parameter gradient and is not checked. + with torch.no_grad(): + m.weight.copy_(w) + return m(x) + + return f"TE {_version('transformer_engine')} RMSNorm(zero_centered_gamma)", fn, "x-only" + + def megatron(): + from megatron.core.transformer.custom_layers.batch_invariant_kernels import ( + BatchInvariantRMSNormFn, + ) + + def fn(x, w): + return BatchInvariantRMSNormFn.apply(x, w, EPS, True) + + # Some revisions accept the flag but still multiply by the uncentered weight. + with torch.no_grad(): + probe_x = torch.ones(2, hidden, device=DEV, dtype=DT) + probe_w = torch.zeros(hidden, device=DEV, dtype=DT) + expected = (probe_x.float() * (1.0 + EPS) ** -0.5).to(DT) + if not torch.allclose(fn(probe_x, probe_w), expected): + raise RuntimeError( + "Megatron zero_centered_gamma=True failed the zero-weight probe; " + "excluded from the zero-centered comparison" + ) + + return "Megatron BatchInvariantRMSNormFn(zero_centered_gamma=True)", fn, "full" + + def flashinfer(): + import flashinfer as fi + + def fn(x, w): + return fi.gemma_rmsnorm(x, w, EPS) + + return f"FlashInfer {_version('flashinfer-python')} gemma_rmsnorm", fn, "none" + + def vllm(): + from vllm.model_executor.layers.layernorm import GemmaRMSNorm + + m = _vllm_module(lambda: GemmaRMSNorm(hidden, eps=EPS)).to(DEV, DT) + + def fn(x, w): + m.weight.data.copy_(w) + return m.forward_cuda(x) + + return f"vLLM {_version('vllm')} GemmaRMSNorm.forward_cuda", fn, "none" + + return [ + ("rl_kernel", ours), + ("reference", reference), + ("transformers", transformers), + ("liger", liger), + ("fla", fla), + ("transformer_engine", te), + ("megatron", megatron), + ("flashinfer", flashinfer), + ("vllm", vllm), + ] + + +def _gated_candidates(hidden: int) -> list[tuple[str, Any]]: + def ours(): + from rl_engine.backends.cuda.norm.rmsnorm import Qwen3NextRMSNormGatedCudaOp + + o = Qwen3NextRMSNormGatedCudaOp() + return "rl-kernel Qwen3NextRMSNormGatedCudaOp", lambda x, z, w: o(x, w, z, eps=EPS), "full" + + def reference(): + from rl_engine.reference.norm.qwen3_next_rms_norm import Qwen3NextRMSNormGatedOp + + o = Qwen3NextRMSNormGatedOp() + return "rl-kernel PyTorch reference", lambda x, z, w: o(x, w, z, eps=EPS), "full" + + def transformers(): + from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextRMSNormGated + + m = Qwen3NextRMSNormGated(hidden, eps=EPS).to(DEV, DT) + + def fn(x, z, w): + return torch.func.functional_call(m, {"weight": w}, (x, z)) + + return f"transformers {_version('transformers')} Qwen3NextRMSNormGated", fn, "full" + + def fla_lg(): + from fla.modules.layernorm_gated import rmsnorm_fn + + def fn(x, z, w): + return rmsnorm_fn(x, w, None, z=z, eps=EPS, group_size=None, norm_before_gate=True) + + return f"FLA {_version('fla-core')} layernorm_gated.rmsnorm_fn", fn, "full" + + def fla_fused(): + from fla.modules.fused_norm_gate import rms_norm_gated + + def fn(x, z, w): + return rms_norm_gated(x, z, w, None, activation="silu", eps=EPS) + + return f"FLA {_version('fla-core')} fused_norm_gate.rms_norm_gated", fn, "full" + + def vllm(): + from vllm.model_executor.layers.layernorm import RMSNormGated + + def build(): + return RMSNormGated( + hidden, eps=EPS, group_size=None, norm_before_gate=True, activation="silu" + ) + + m = _vllm_module(build).to(DEV, DT) + + def fn(x, z, w): + m.weight.data.copy_(w) + return m.forward_cuda(x, z) + + return f"vLLM {_version('vllm')} RMSNormGated.forward_cuda", fn, "none" + + return [ + ("rl_kernel", ours), + ("reference", reference), + ("transformers", transformers), + ("fla_layernorm_gated", fla_lg), + ("fla_fused_norm_gate", fla_fused), + ("vllm", vllm), + ] + + +def candidates(op: str) -> dict[str, dict[str, Any]]: + hidden = SHAPES[op][0] + found = _c1_candidates(hidden) if op == "qwen3_next_rms_norm" else _gated_candidates(hidden) + out: dict[str, dict[str, Any]] = {} + for name, factory in found: + try: + source, fn, backward = factory() + out[name] = {"source": source, "fn": fn, "backward": backward} + except Exception as exc: # noqa: BLE001 - optional library missing or unusable + if name == "rl_kernel": + raise RuntimeError(f"Required rl_kernel candidate is unavailable: {exc}") from exc + out[name] = {"unavailable": f"{type(exc).__name__}: {exc}"[:300]} + return out + + +def _inputs(op: str, seed: int, n: int) -> dict[str, torch.Tensor]: + hidden = SHAPES[op][0] + g = torch.Generator(device=DEV).manual_seed(seed) + x = (torch.randn(n, hidden, device=DEV, generator=g) * 2).to(DT) + if op == "qwen3_next_rms_norm": + return {"x": x, "w": (torch.randn(hidden, device=DEV, generator=g) * 0.1).to(DT)} + z = torch.randn(n, hidden, device=DEV, generator=g).to(DT) + return {"x": x, "z": z, "w": (1 + torch.randn(hidden, device=DEV, generator=g) * 0.1).to(DT)} + + +def _reference(op: str, inputs: dict[str, torch.Tensor]): + x = inputs["x"].double() + n = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) + if op == "qwen3_next_rms_norm": + return n * (1 + inputs["w"].double()) + return n * inputs["w"].double() * F.silu(inputs["z"].double()) + + +# --------------------------------------------------------------------------- # +# Batch invariance +# --------------------------------------------------------------------------- # + + +def _run(fn, inputs, rows, grads, ups, idx=None): + args = {} + for k, v in inputs.items(): + if idx is not None and k in rows: + v = v.index_select(0, idx) + args[k] = v.detach().clone().requires_grad_(True) if k in grads else v + out = fn(**args) + if grads: + out.backward(ups if idx is None else ups.index_select(0, idx)) + return out.detach(), {k: args[k].grad for k in grads} + + +def _sweep_sizes(n: int) -> list[int]: + sizes = set(range(1, min(n, 9) + 1)) + k = 16 + while k <= n: + sizes.update(v for v in (k - 1, k, k + 1) if v <= n) + k *= 2 + return sorted(sizes | {n}) + + +def batch_invariance(op: str, cand: dict[str, Any], quick: bool) -> dict[str, Any]: + _, full_sizes, big, subs, _ = SHAPES[op] + if quick: + full_sizes, big, subs = (17, 64), 512, (1, 7, 64) + fn = cand["fn"] + rows = ["x"] if op == "qwen3_next_rms_norm" else ["x", "z"] + grads = {"full": rows + ["w"], "x-only": rows, "none": []}[cand["backward"]] + row_grads = [k for k in rows if k in grads] + + def ups_for(seed, n): + g = torch.Generator(device=DEV).manual_seed(seed + 100) + return torch.randn(n, SHAPES[op][0], device=DEV, generator=g).to(DT) + + res: dict[str, Any] = {"all_rows": {}, "dweight_repeatable": None} + ok = True + for n in full_sizes: + fwd = grd = 0 + repeat = True + for seed in (3, 4, 5): + inp, ups = _inputs(op, seed, n), ups_for(seed, n) + fo, fg = _run(fn, inp, rows, grads, ups) + for i in range(n): + o, g = _run(fn, inp, rows, grads, ups, torch.tensor([i], device=DEV)) + fwd += not torch.equal(o[0], fo[i]) + grd += not all(torch.equal(g[k][0], fg[k][i]) for k in row_grads) + if "w" in grads: + _, again = _run(fn, inp, rows, grads, ups) + repeat &= torch.equal(again["w"], fg["w"]) + if "w" in grads: + res["dweight_repeatable"] = repeat and res["dweight_repeatable"] is not False + res["all_rows"][str(n)] = {"rows": 3 * n, "fwd_differ": fwd, "rowgrad_differ": grd} + ok &= fwd == 0 and grd == 0 and repeat + + fwd = grd = checked = 0 + for seed in (3, 4): + inp, ups = _inputs(op, seed, big), ups_for(seed, big) + fo, fg = _run(fn, inp, rows, grads, ups) + for sb in subs: + for start in range(0, big, sb): + idx = torch.arange(start, min(start + sb, big), device=DEV) + o, g = _run(fn, inp, rows, grads, ups, idx) + checked += 1 + fwd += not torch.equal(o, fo.index_select(0, idx)) + grd += not all(torch.equal(g[k], fg[k].index_select(0, idx)) for k in row_grads) + del inp, ups, fo, fg + torch.cuda.empty_cache() + res["full_vs_sub_batches"] = { + "full_batch": big, + "sub_batch_sizes": list(subs), + "coverage": "every_row_per_sub_batch_size", + "sub_batches": checked, + "fwd_differ": fwd, + "rowgrad_differ": grd, + } + ok &= fwd == 0 and grd == 0 + + n = full_sizes[-1] + sizes = _sweep_sizes(n) + probes = sorted({0, n - 1, *list(range(n // 8 + 3, n - 1, n // 8))[:6]}) + compared = failed = 0 + for seed in (3, 4, 5): + inp, ups = _inputs(op, seed, n), ups_for(seed, n) + alone = {r: _run(fn, inp, rows, grads, ups, torch.tensor([r], device=DEV)) for r in probes} + gen = torch.Generator().manual_seed(seed) + for m in sizes: + for r in probes: + others = torch.randperm(n, generator=gen) + others = others[others != r][: m - 1] + for pos in sorted({0, (m - 1) // 2, m - 1}): + idx = torch.cat([others[:pos], torch.tensor([r]), others[pos:]]).to(DEV) + o, g = _run(fn, inp, rows, grads, ups, idx) + good = torch.equal(o[pos], alone[r][0][0]) + good &= all(torch.equal(g[k][pos], alone[r][1][k][0]) for k in row_grads) + compared += 1 + failed += not good + res["batch_size_sweep"] = { + "batch_sizes": sizes, + "probe_rows": probes, + "comparisons": compared, + "differ": failed, + } + ok &= failed == 0 + res["batch_invariant"] = bool(ok) + return res + + +# --------------------------------------------------------------------------- # +# Accuracy and latency +# --------------------------------------------------------------------------- # + + +def _time_us(fn, warmup=10, iters=50) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(iters): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1e3) + return statistics.median(samples) + + +def accuracy_latency(op: str, cand: dict[str, Any], quick: bool) -> dict[str, Any]: + n = 4096 if quick else SHAPES[op][4] + inp = _inputs(op, 11, n) + keys = list(inp) + gk = {"full": keys, "x-only": [k for k in keys if k != "w"], "none": []}[cand["backward"]] + fn = cand["fn"] + leaves64 = {k: v.double().requires_grad_(k in gk) for k, v in inp.items()} + ref = _reference(op, leaves64) + with torch.no_grad(): + y = fn(**inp) + out: dict[str, Any] = { + "rows": n, + "fwd_correctly_rounded": (y == ref.detach().to(DT)).float().mean().item(), + "fwd_max_abs_err": (y.double() - ref.detach()).abs().max().item(), + "fwd_us": _time_us(lambda: fn(**inp)), + } + if gk: + gen = torch.Generator(device=DEV).manual_seed(5) + dy = torch.randn(n, SHAPES[op][0], device=DEV, generator=gen).to(DT) + rg = torch.autograd.grad(ref, [leaves64[k] for k in gk], dy.double()) + lv = {k: v.detach().clone().requires_grad_(k in gk) for k, v in inp.items()} + gs = torch.autograd.grad(fn(**lv), [lv[k] for k in gk], dy) + out["grad_rel_err"] = { + k: ((g.double() - r).abs().max() / r.abs().max()).item() for k, g, r in zip(gk, gs, rg) + } + + def step(): + lv = {k: v.detach().clone().requires_grad_(k in gk) for k, v in inp.items()} + torch.autograd.grad(fn(**lv), [lv[k] for k in gk], dy) + + out["fwd_bwd_us"] = _time_us(step) + return out + + +# --------------------------------------------------------------------------- # +# The repository's C3/C4 gates with an existing implementation swapped in +# --------------------------------------------------------------------------- # + + +def _gate_child(which: str, op: str, impl: str) -> None: + """Run the canonical C3/C4 CLI with the CUDA candidate replaced by ``impl``.""" + + import importlib.util + + from rl_engine.backends.cuda.norm import rmsnorm as R + + path = REPO_ROOT / "tools/validation/operators" / f"check_{which}_invariance.py" + spec = importlib.util.spec_from_file_location("gate", path) + gate = importlib.util.module_from_spec(spec) + sys.argv = [ + str(path), + "--manifest", + MANIFEST, + "--op", + op, + "--candidate", + "cuda", + "--backend-profile", + "cuda_bf16", + "--hidden", + "2048", + "--head-dim", + "128", + ] + spec.loader.exec_module(gate) + if impl != "rl_kernel": + fn = candidates(op)[impl]["fn"] + if op == "qwen3_next_rms_norm": + + class Swapped(R.Qwen3NextRMSNormCudaOp): + def forward(self, x, weight, *, eps=EPS): + return fn(x.reshape(-1, x.shape[-1]), weight).view_as(x) + + else: + + class Swapped(R.Qwen3NextRMSNormGatedCudaOp): + def __call__(self, x, weight, gate, *, eps=EPS): + return self.forward(x, weight, gate, eps=eps) + + def forward(self, x, weight, gate, *, eps=EPS): + return fn(x, gate, weight) + + original = gate.load_adapter_operator + + def load_adapter_operator(op_name, candidate): + return Swapped() if candidate == "cuda" else original(op_name, candidate) + + gate.load_adapter_operator = load_adapter_operator + gate.main() + + +def _clean(lines: list[str], megatron_src: str | None) -> list[str]: + """Drop warnings and replace machine-specific paths, so reports carry no local paths.""" + + subs = [(str(REPO_ROOT), ""), (sys.prefix, ""), (str(Path.home()), "")] + if megatron_src: + subs.insert(0, (str(Path(megatron_src).resolve()), "")) + out = [] + for line in lines: + if "Warning" in line or "warnings.warn" in line: + continue + for path, name in subs: + line = line.replace(path, name) + out.append(line) + return out + + +def contract_gates(op: str, impl: str, megatron_src: str | None) -> dict[str, Any]: + if not (REPO_ROOT / MANIFEST).exists(): + # The Qwen3-Next C3/C4 gate adapters and manifest arrive with the gated-norm PR. + return {"unavailable": f"{MANIFEST} is not on this branch"} + out: dict[str, Any] = {} + for which in ("forward", "gradient"): + cmd = [ + sys.executable, + __file__, + "--op", + op, + "--out", + os.devnull, + "--gate-child", + which, + impl, + ] + if megatron_src: + cmd += ["--megatron-src", megatron_src] + proc = subprocess.run( + cmd, capture_output=True, text=True, env={**os.environ, "RL_KERNEL_REQUIRE_EXT": "1"} + ) + lines = _clean(proc.stdout.splitlines(), megatron_src) + summary = next((ln for ln in lines if ln.startswith("op=")), "") + out[which] = { + "returncode": proc.returncode, + "passed": "passed=True" in summary and proc.returncode == 0, + "failed_lines": [ln.strip() for ln in lines if "passed=False" in ln][:10], + "singleton_aggregate": [ln.strip() for ln in lines if "singleton_aggregate pair" in ln], + "stderr_tail": ( + _clean(proc.stderr.strip().splitlines(), megatron_src)[-3:] + if proc.returncode + else [] + ), + } + return out + + +# --------------------------------------------------------------------------- # +# Report and figure +# --------------------------------------------------------------------------- # + + +def plot(report: dict[str, Any], path: Path) -> None: + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + except ImportError: + return + entries = [(n, e) for n, e in report["results"].items() if "unavailable" not in e] + names = [e["source"] for _, e in entries] + ys = list(range(len(entries))) + fig, axes = plt.subplots(1, 3, figsize=(22, 0.55 * len(entries) + 3), layout="constrained") + fig.suptitle( + f"{report['op']} vs existing implementations — {report['environment']['gpu']}, " + f"commit {report['rl_kernel_commit'][:7]}" + ) + ax = axes[0] + for key, off, label in (("fwd_us", -0.2, "forward"), ("fwd_bwd_us", 0.2, "fwd + bwd")): + vals = [e.get("accuracy_latency", {}).get(key, float("nan")) for _, e in entries] + ax.barh([y + off for y in ys], vals, 0.4, label=label) + ax.set_xscale("log") + ax.set_yticks(ys, names, fontsize=7) + ax.invert_yaxis() + ax.set_xlabel("µs, median") + ax.set_title("latency at the workload size") + ax.legend(fontsize=7) + ax.grid(True, axis="x", alpha=0.3) + ax = axes[1] + acc = [e.get("accuracy_latency", {}) for _, e in entries] + cr = [a.get("fwd_correctly_rounded", float("nan")) for a in acc] + ax.barh(ys, [max(1 - c, 1e-7) for c in cr]) + for y, c, a in zip(ys, cr, acc): + grads = a.get("grad_rel_err", {}) + worst = f"; worst grad err {max(grads.values()):.1e}" if grads else "" + ax.text(1.5e-7, y, f"{c:.4%} correctly rounded{worst}", va="center", fontsize=7) + ax.set_xscale("log") + ax.set_xlim(1e-7, 1) + ax.set_yticks(ys, [""] * len(entries)) + ax.invert_yaxis() + ax.set_title("forward: fraction NOT equal to the correctly rounded FP64 result") + ax = axes[2] + for y, (_, e) in zip(ys, entries): + bi = e.get("batch_invariance") + gates = e.get("contract_gates") + text = "checks passed: " + ( + "n/a" if bi is None else ("yes" if bi["batch_invariant"] else "NO") + ) + if bi and bi["full_vs_sub_batches"].get("coverage") == "sampled_small_sub_batches": + text += " (sampled)" + if gates and "unavailable" not in gates: + text += " | C3/C4 gates: " + ( + "pass" if all(g["passed"] for g in gates.values()) else "FAIL" + ) + good = bi is not None and bi["batch_invariant"] + ax.text(0.02, y, text, va="center", fontsize=8, color="#3aa676" if good else "#d14b4b") + ax.set_ylim(len(entries) - 0.5, -0.5) + ax.axis("off") + ax.set_title("measured batch-invariance checks and this repository's gates") + fig.savefig(path.with_suffix(".png"), dpi=120) + + +def _git(*args: str) -> str: + return subprocess.run( + ["git", *args], cwd=REPO_ROOT, capture_output=True, text=True + ).stdout.strip() + + +def main() -> None: + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument("--op", required=True, choices=sorted(SHAPES)) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--checks", default="bi,perf,gates") + parser.add_argument("--megatron-src", default=None, help="Megatron-LM source tree (optional)") + parser.add_argument("--quick", action="store_true", help="small sizes, for a smoke test only") + parser.add_argument("--gate-child", nargs=2, default=None, help=argparse.SUPPRESS) + args = parser.parse_args() + if args.megatron_src: + sys.path.insert(0, args.megatron_src) + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + if args.gate_child: + _gate_child(args.gate_child[0], args.op, args.gate_child[1]) + return + if args.op == "rms_norm_gated": + from rl_engine.backends.cuda.norm import rmsnorm as R + + if not hasattr(R, "Qwen3NextRMSNormGatedCudaOp"): + raise SystemExit("rms_norm_gated: the gated CUDA op is not on this branch") + torch.backends.cuda.matmul.allow_tf32 = False + checks = set(args.checks.split(",")) + results: dict[str, Any] = {} + for name, cand in candidates(args.op).items(): + if "unavailable" in cand: + results[name] = cand + print(f"{name}: unavailable ({cand['unavailable'][:80]})", flush=True) + continue + entry: dict[str, Any] = {"source": cand["source"], "backward": cand["backward"]} + if "bi" in checks: + entry["batch_invariance"] = batch_invariance(args.op, cand, args.quick) + if "perf" in checks: + entry["accuracy_latency"] = accuracy_latency(args.op, cand, args.quick) + if "gates" in checks and cand["backward"] == "full" and name != "reference": + entry["contract_gates"] = contract_gates(args.op, name, args.megatron_src) + results[name] = entry + summary = {k: v for k, v in entry.items() if k != "source"} + print(f"{name}: {json.dumps(summary)[:300]}", flush=True) + torch.cuda.empty_cache() + + megatron_commit = None + if args.megatron_src: + megatron_commit = ( + subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=args.megatron_src, + capture_output=True, + text=True, + ).stdout.strip() + or None + ) + report = { + "kind": "qwen3_next_norm_reuse_check", + "rfc": "RL-Align/RL-Kernel#428", + "op": args.op, + "rl_kernel_commit": _git("rev-parse", "HEAD"), + "tracked_tree_dirty": bool(_git("status", "--porcelain", "--untracked-files=no")), + "environment": { + "gpu": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "cuda": torch.version.cuda, + "libraries": {lib: _version(lib) for lib in LIBS}, + "megatron_source_commit": megatron_commit, + }, + "quick": args.quick, + "checks": sorted(checks), + "results": results, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + plot(report, args.out) + commit, dirty = report["rl_kernel_commit"][:7], report["tracked_tree_dirty"] + print(f"wrote {args.out} (commit {commit}, dirty={dirty})") + + +if __name__ == "__main__": + main() 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: