Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
22f2ac3
fix(registry): dispatch rms_norm to the CUDA backend on CUDA
fusheng-ji Sep 30, 2026
ba9b89e
fix(norm): validate only used CUDA symbols and exercise explicit fall…
fusheng-ji Sep 30, 2026
1dd6dbb
test(rmsnorm): gate the dispatch test on the symbols the CUDA op requ…
fusheng-ji Oct 4, 2026
6a4f078
fix(rmsnorm): guard the CUDA launchers on the input's device
fusheng-ji Oct 4, 2026
6741d16
feat(ws1): Qwen3-Next RMSNorm references and zero-centred CUDA offset
fusheng-ji Sep 30, 2026
5f5654a
fix(norm): preserve signed zero and use CUDA statistics for parameter…
fusheng-ji Sep 30, 2026
d03c852
fix(norm): contract tolerances and corrected claims for the Qwen3-Nex…
fusheng-ji Oct 1, 2026
81eae9e
Merge test-qwennext into Qwen3-Next C1 norm branch
fusheng-ji Oct 6, 2026
bec7e7a
fix(norm): reject stale RMSNorm bindings and enable GPU coverage
fusheng-ji Oct 7, 2026
53767ab
feat(scripts): Qwen3-Next norm evidence runner vs existing implementa…
fusheng-ji Oct 8, 2026
3d0bae7
fix(scripts): readable Qwen3-Next norm evidence figure layout
fusheng-ji Oct 8, 2026
ba46750
docs(norm): zero-centred RMSNorm evidence vs existing implementations…
fusheng-ji Oct 8, 2026
904d6fc
feat(scripts): Qwen3-Next norm reuse check against existing implement…
fusheng-ji Oct 8, 2026
791e40d
docs(norm): qwen3_next_rms_norm vs existing implementations, every-ro…
fusheng-ji Oct 8, 2026
af13778
fix(scripts): keep local paths and warnings out of norm reuse-check g…
fusheng-ji Oct 8, 2026
e39f307
fix: exclude invalid Megatron zero-centered norm evidence
fusheng-ji Oct 9, 2026
19b7dc3
Merge test-qwennext refactor and migrate Qwen3-Next code and tests
fusheng-ji Oct 10, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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: |
Expand Down
6 changes: 4 additions & 2 deletions .github/workflows/ws1-gtest-gpu.yml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ name: WS1-gtest-GPU

on:
pull_request:
branches: [ main ]
branches: [ main, test-qwennext ]
paths:
- "rl_engine/validation/**"
- "rl_engine/backends/**"
Expand All @@ -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"
Expand All @@ -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/**"
Expand All @@ -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"
Expand Down
1 change: 1 addition & 0 deletions ci/scripts/run_ws1_gtest.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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 \
Expand Down
26 changes: 18 additions & 8 deletions csrc/bindings/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -263,14 +263,16 @@ void rmsnorm_forward_cuda(
torch::Tensor weight,
torch::Tensor y,
torch::Tensor rstd,
double eps);
double eps,
double weight_offset);

void rmsnorm_backward_dx_cuda(
torch::Tensor dy,
torch::Tensor x,
torch::Tensor weight,
torch::Tensor rstd,
torch::Tensor dx);
torch::Tensor dx,
double weight_offset);

void rmsnorm_backward_partial_dw_cuda(
torch::Tensor dy,
Expand Down Expand Up @@ -299,7 +301,8 @@ static void rmsnorm_check_input(const torch::Tensor& x, const char* name) {
std::vector<torch::Tensor> 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");
Expand All @@ -312,7 +315,7 @@ std::vector<torch::Tensor> rmsnorm_forward(
auto y = torch::empty_like(x);
auto rstd = torch::empty({T}, x.options().dtype(torch::kFloat32));

rmsnorm_forward_cuda(x, weight, y, rstd, eps);
rmsnorm_forward_cuda(x, weight, y, rstd, eps, weight_offset);

return {y, rstd};
}
Expand All @@ -321,7 +324,8 @@ torch::Tensor rmsnorm_backward_dx(
torch::Tensor dy,
torch::Tensor x,
torch::Tensor weight,
torch::Tensor rstd)
torch::Tensor rstd,
double weight_offset)
{
rmsnorm_check_input(dy, "dy");
rmsnorm_check_input(x, "x");
Expand All @@ -336,7 +340,7 @@ torch::Tensor rmsnorm_backward_dx(

auto dx = torch::empty_like(x);

rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx);
rmsnorm_backward_dx_cuda(dy, x, weight, rstd, dx, weight_offset);

return dx;
}
Expand Down Expand Up @@ -718,8 +722,14 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
&det_gemm_db_transposed,
"Batch-invariant deterministic GEMM backward in canonical [N,K] layout");
// registry RMSNorm
m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA");
m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA");
// API version 2 includes weight_offset in both entry points.
m.attr("rmsnorm_api_version") = 2;
m.def("rmsnorm_forward", &rmsnorm_forward, "Batch-invariant RMSNorm forward CUDA",
py::arg("x"), py::arg("weight"), py::arg("eps"),
py::arg("weight_offset") = 0.0);
m.def("rmsnorm_backward_dx", &rmsnorm_backward_dx, "Batch-invariant RMSNorm backward dx CUDA",
py::arg("dy"), py::arg("x"), py::arg("weight"), py::arg("rstd"),
py::arg("weight_offset") = 0.0);
m.def("rmsnorm_backward_dw", &rmsnorm_backward_dw, "Deterministic RMSNorm backward dweight CUDA");
#if !defined(USE_ROCM)
m.def(
Expand Down
29 changes: 22 additions & 7 deletions csrc/cuda/norm/rmsnorm.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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<scalar_t>(x_row + col);
float wv = load_as_float<weight_t>(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<scalar_t>(y_row + col, out);
}
Expand All @@ -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;
Expand All @@ -168,6 +177,7 @@ __global__ void rmsnorm_bwd_dx_kernel(
float dyv = load_as_float<scalar_t>(dy_row + col);
float xv = load_as_float<scalar_t>(x_row + col);
float wv = load_as_float<weight_t>(weight + col);
if (weight_offset != 0.0f) wv += weight_offset;
local_dot += dyv * wv * xv;
}

Expand All @@ -180,6 +190,7 @@ __global__ void rmsnorm_bwd_dx_kernel(
float dyv = load_as_float<scalar_t>(dy_row + col);
float xv = load_as_float<scalar_t>(x_row + col);
float wv = load_as_float<weight_t>(weight + col);
if (weight_offset != 0.0f) wv += weight_offset;

float out = r * dyv * wv - xv * coeff;
store_from_float<scalar_t>(dx_row + col, out);
Expand Down Expand Up @@ -247,7 +258,8 @@ void rmsnorm_forward_cuda(
torch::Tensor weight,
torch::Tensor y,
torch::Tensor rstd,
double eps
double eps,
double weight_offset
) {
// Launch on x's device: the current CUDA stream belongs to the current
// device, which need not be x's.
Expand All @@ -270,7 +282,8 @@ void rmsnorm_forward_cuda(
rstd.data_ptr<float>(),
T,
H,
static_cast<float>(eps)
static_cast<float>(eps),
static_cast<float>(weight_offset)
);
});
});
Expand All @@ -282,7 +295,8 @@ void rmsnorm_backward_dx_cuda(
torch::Tensor x,
torch::Tensor weight,
torch::Tensor rstd,
torch::Tensor dx
torch::Tensor dx,
double weight_offset
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
int T = x.size(0);
Expand All @@ -303,7 +317,8 @@ void rmsnorm_backward_dx_cuda(
rstd.data_ptr<float>(),
dx.data_ptr<x_t>(),
T,
H
H,
static_cast<float>(weight_offset)
);
});
});
Expand Down
1 change: 1 addition & 0 deletions docs/.nav.yml
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ nav:
- operators/sampling.md
- operators/det-gemm.md
- operators/embedding.md
- operators/qwen3-next-rms-norm.md
- Developer Guide:
- contributing/README.md
- Contributor Guide: contributing/contributor-guide.md
Expand Down
1 change: 1 addition & 0 deletions docs/operators/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -32,4 +32,5 @@ Every operator page should include:
- [Matmul](matmul.md)
- [Sampling](sampling.md)
- [Token Embedding](embedding.md)
- [Qwen3-Next RMSNorm (zero-centred)](qwen3-next-rms-norm.md)
- [Operator Doc Template](../contributing/operator-doc-template.md)
Loading
Loading