Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
39 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
b11b957
feat(cuda): gated RMSNorm kernel for the Qwen3-Next GDN block
fusheng-ji Sep 30, 2026
d03c852
fix(norm): contract tolerances and corrected claims for the Qwen3-Nex…
fusheng-ji Oct 1, 2026
7dd291a
fix(norm): enforce gated tensor boundaries and align parameter VJP st…
fusheng-ji Sep 30, 2026
8df1ffb
test(norm): add Qwen3-Next workload and C3 C4 adapters
fusheng-ji Sep 30, 2026
aaf0986
fix(norm): preserve gated signed-zero weights when offset is disabled
fusheng-ji Sep 30, 2026
4ea1611
test(dispatch): cover qwen3_next_rms_norm priority and missing-extens…
fusheng-ji Oct 1, 2026
5eaeca8
test(norm): sweep the gated-vs-plain rstd identity across dtype, H, a…
fusheng-ji Oct 1, 2026
b0f4c5c
fix(norm): address the gated-norm review (R2-2, 7, 8, 9, 12, 13, 14, …
fusheng-ji Oct 1, 2026
0cceb84
fix(rmsnorm): use OptionalCUDAGuard in the launchers outside the ROCm…
fusheng-ji Oct 2, 2026
cf68fd9
docs(norm): reconcile the C1 norm page, docstrings and gated tests wi…
fusheng-ji Oct 4, 2026
81eae9e
Merge test-qwennext into Qwen3-Next C1 norm branch
fusheng-ji Oct 6, 2026
8ce67e0
Merge updated C1 norm base into gated RMSNorm 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
1edc9d4
Merge Qwen3-Next C1 norm fixes and evidence runner into gated RMSNorm…
fusheng-ji Oct 8, 2026
a441ff9
feat(scripts): extend the Qwen3-Next norm evidence runner to the gate…
fusheng-ji Oct 8, 2026
3d0bae7
fix(scripts): readable Qwen3-Next norm evidence figure layout
fusheng-ji Oct 8, 2026
822b085
Merge Qwen3-Next C1 norm plot fix into gated RMSNorm branch
fusheng-ji Oct 8, 2026
ba46750
docs(norm): zero-centred RMSNorm evidence vs existing implementations…
fusheng-ji Oct 8, 2026
024823d
docs(norm): gated RMSNorm evidence vs existing implementations on B200
fusheng-ji Oct 8, 2026
062fc5a
Merge Qwen3-Next C1 norm evidence into gated RMSNorm branch
fusheng-ji Oct 8, 2026
b4bb977
docs(norm): correct the gated backward slowdown ratio
fusheng-ji Oct 8, 2026
904d6fc
feat(scripts): Qwen3-Next norm reuse check against existing implement…
fusheng-ji Oct 8, 2026
08de15b
Merge the Qwen3-Next norm reuse check into PR #468
fusheng-ji Oct 8, 2026
791e40d
docs(norm): qwen3_next_rms_norm vs existing implementations, every-ro…
fusheng-ji Oct 8, 2026
bb4ff71
Merge the #467 norm reuse evidence into PR #468
fusheng-ji Oct 8, 2026
af13778
fix(scripts): keep local paths and warnings out of norm reuse-check g…
fusheng-ji Oct 8, 2026
a8d5fa5
Merge the norm reuse-check output fix into PR #468
fusheng-ji Oct 8, 2026
4783c1c
docs(norm): C3/C4 gates for the zero-centred norm with each existing …
fusheng-ji Oct 8, 2026
eeb8b0a
docs(norm): rms_norm_gated vs existing implementations, every-row bat…
fusheng-ji Oct 8, 2026
955ae3a
fix: exclude invalid Megatron zero-centered norm evidence
fusheng-ji Oct 9, 2026
8e042fd
fix: make norm reuse checks exhaustive and require primary candidate
fusheng-ji Oct 9, 2026
910bee3
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
12 changes: 12 additions & 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 Expand Up @@ -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
151 changes: 137 additions & 14 deletions csrc/bindings/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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<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");

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};
}
Expand All @@ -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<torch::Tensor> 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;
}
Expand All @@ -350,13 +455,16 @@ 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");
TORCH_CHECK(x.dim() == 2, "x must be 2D [T, H]");
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)");

Expand Down Expand Up @@ -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",
Expand Down
Loading
Loading