Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
71 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
ef4c88d
feat(ws1): Gated DeltaNet decode-step goldens for RFC #428 C6
fusheng-ji Sep 30, 2026
dbb8d00
fix(gdn): validate cache ownership and reproduce conv product rounding
fusheng-ji Sep 30, 2026
dcb9263
fix(gdn): avoid overflow in inactive softplus gradient branch
fusheng-ji Sep 30, 2026
16204ab
ci(qwen3-next): separate audited provider runtime from legacy CUDA gate
fusheng-ji Sep 30, 2026
697bea8
test(gdn): assert the packed-decode provider leaves the null block un…
fusheng-ji Oct 1, 2026
12a2225
ci(qwen3-next): drop stale and forward-referencing workflow paths
fusheng-ji Oct 1, 2026
b54ac72
feat(gdn): add a runner for the C6 provider-agreement measurements
fusheng-ji Oct 1, 2026
7b30633
docs(gdn): correct the C6 design note against vLLM 0.30.0 source
fusheng-ji Oct 1, 2026
9da9c36
fix(gdn): report unknown git state as null, not dirty, in the C6 runner
fusheng-ji Oct 1, 2026
5216a6c
docs(gdn): restate the backward deferral against RFC #428
fusheng-ji Oct 1, 2026
10e5230
fix(gdn): make the C6 runner's fusion-off arm and FMA count real
fusheng-ji Oct 1, 2026
2d895d0
docs(gdn): quote the runner's measured C6 agreement, not stale counts
fusheng-ji Oct 1, 2026
52f7f61
docs(gdn): record that FP contraction does not explain the conv misma…
fusheng-ji Oct 1, 2026
7862a12
docs(gdn): scope the mixed-step caveat to what was measured
fusheng-ji Oct 1, 2026
1d69312
fix(gdn): address the C6 review (R3 F1-4, F6, F10-14, F17, F20)
fusheng-ji Oct 1, 2026
c0ab21c
docs(gdn): refresh the design note's in-repo line references
fusheng-ji Oct 1, 2026
f8d9399
feat(gdn): localise the fp32-cache conv mismatches in the C6 runner
fusheng-ji Oct 2, 2026
610f1b7
docs(gdn): state the conv mismatch causes, measured by the runner
fusheng-ji Oct 2, 2026
395fef8
ci(qwen3-next): reconcile the WS1 gtest script with the stacked gated…
fusheng-ji Oct 4, 2026
0b41350
fix(gdn): keep _chunked_sum fixed-order for widths not divisible by 32
fusheng-ji Oct 4, 2026
c883134
Merge test-qwennext into feat/ws1-c6-gdn-recurrent-golden
fusheng-ji Oct 6, 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
490499c
Merge aligned gated RMSNorm branch (C1 fixes, evidence runner) into C…
fusheng-ji Oct 8, 2026
db420ab
test(gdn): L1 batch-invariance check for the causal-conv1d golden
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
c214004
Merge gated RMSNorm evidence into C6 GDN goldens
fusheng-ji Oct 8, 2026
2e950f3
feat(qwen3-next): MoE route/combine contract with a batch-invariant g…
fusheng-ji Oct 9, 2026
9e6c4f2
feat(qwen3-next): TP4 GDN, convolution and D=256 attention mixers
fusheng-ji Oct 9, 2026
3fe632d
docs(qwen3-next): MoE route/combine contract, prior-art and TP4 evide…
fusheng-ji Oct 9, 2026
811a338
Merge MoE route/combine evidence into the TP4 mixers
fusheng-ji Oct 9, 2026
6517077
test(qwen3-next): prior-art runner for the TP4 D=256 attention core
fusheng-ji Oct 9, 2026
b6fddb3
docs(qwen3-next): TP4 mixer gates and D=256 attention prior-art evide…
fusheng-ji Oct 9, 2026
c58816b
test(qwen3-next): measure Megatron-core/TE and SGLang in the MoE prio…
fusheng-ji Oct 9, 2026
27b7339
docs(qwen3-next): MoE prior art with Megatron-core/TE and SGLang meas…
fusheng-ji Oct 9, 2026
5bbb2bc
Merge the MoE prior-art update into the TP4 mixers
fusheng-ji Oct 10, 2026
4ef7dcb
test(qwen3-next): measure FlashAttention-4, Transformer Engine and Me…
fusheng-ji Oct 10, 2026
703c6dc
fix(qwen3-next): merge flat and per-size attention latency sections
fusheng-ji Oct 10, 2026
414a53e
docs(qwen3-next): attention prior art with FlashAttention-4, TE and M…
fusheng-ji Oct 10, 2026
7f70b60
Merge test-qwennext refactor and migrate Qwen3-Next code and tests
fusheng-ji Oct 10, 2026
9496e72
Merge test-qwennext refactor and migrate Qwen3-Next code and tests
fusheng-ji Oct 10, 2026
1513b71
style(tools): mark the Qwen3-Next TP gate imports after sys.path setu…
fusheng-ji Oct 10, 2026
69d015a
style(tools): mark the Qwen3-Next TP gate imports after sys.path setu…
fusheng-ji Oct 10, 2026
37bc953
fix(qwen3-next): type annotations for mypy in the MoE prior-art runner
fusheng-ji Oct 11, 2026
f861fbe
Merge the MoE prior-art type fix into the TP4 mixers
fusheng-ji Oct 11, 2026
fac8f92
fix(qwen3-next): type annotations for mypy in the attention prior-art…
fusheng-ji Oct 11, 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 tests/models/qwen3_next/test_gdn_state_contract.py tests/validation/common/test_tensor_identity.py tests/models/qwen3_next/test_qwen3_next_forward_contract.py tests/models/qwen3_next/test_qwen3_next_tp_blocks.py tests/models/qwen3_next/test_qwen3_next_tp_gdn.py -q

- name: Run Cross-Configuration Contract Tests (CPU-safe)
run: |
Expand Down
95 changes: 95 additions & 0 deletions .github/workflows/qwen3-next-provider-gpu.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
# RFC #428 Qwen3-Next provider comparisons: the C1/C3/C4 norm providers, the C6
# GDN decode-step goldens, the shared GEMM/MoE route-combine provider and the TP4
# GDN/convolution/attention mixers, each checked against the real vLLM 0.30.0 provider.
#
# These check_ files import vLLM, so the default pytest collection never runs them
# (tests/integrations/common/test_framework_operator_integrations.py asserts vLLM is not imported). This
# workflow runs them in their own process on a self-hosted runner that a maintainer
# registers with the labels below and keeps on the audited runtime (CUDA 13,
# torch 2.13.0+cu130, vLLM 0.30.0). Without such a runner the job queues rather than
# reporting a false pass.
#
# Security: do not use pull_request_target. Fork PRs never reach the self-hosted
# runner; a maintainer dispatches the reviewed commit from a trusted branch.

name: Qwen3-Next-provider-GPU

on:
pull_request:
branches: [main, test-qwennext]
paths:
- "rl_engine/backends/**"
- "rl_engine/reference/**"
- "csrc/cuda/attention/**"
- "rl_engine/models/qwen3_next/**"
- "rl_engine/integrations/engines/train/vllm/qwen3_next*"
- "rl_engine/validation/models/qwen3_next*"
- "tools/validation/models/*qwen3_next*"
- "tests/models/qwen3_next/**"
- "tests/models/qwen3_next/check_gdn_recurrent_golden.py"
- ".github/workflows/qwen3-next-provider-gpu.yml"
workflow_dispatch:

concurrency:
group: qwen3-next-provider-gpu-${{ github.ref }}
cancel-in-progress: false

permissions:
contents: read

jobs:
fork-pr-notice:
if: github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository
runs-on: ubuntu-latest
steps:
- name: Report required trusted execution
run: |
echo "Fork code does not run on the self-hosted Qwen3-Next provider runner."
echo "A maintainer must dispatch this workflow from a trusted upstream branch."
echo "source_repository=${{ github.event.pull_request.head.repo.full_name }}"
echo "source_sha=${{ github.event.pull_request.head.sha }}"

provider:
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository
runs-on: [self-hosted, linux, x64, rl-kernel-qwen3-next]
timeout-minutes: 30
steps:
- uses: actions/checkout@v4
with:
persist-credentials: false
- name: Require the audited runtime
run: |
python3 - <<'PY'
import torch, vllm
assert torch.__version__ == "2.13.0+cu130"
assert vllm.__version__ == "0.30.0"
assert torch.cuda.is_available()
PY
- name: Build the extension against that runtime
run: python3 setup.py build_ext --inplace
- name: Require the extension built from this checkout
# A stale editable install elsewhere on the runner would otherwise satisfy the
# import and test some other tree's _C.
run: |
python3 - <<'PY'
import os
from pathlib import Path

from rl_engine import _C

workspace = Path(os.environ["GITHUB_WORKSPACE"]).resolve()
loaded = Path(_C.__file__).resolve()
assert workspace in loaded.parents, f"rl_engine._C loaded from {loaded}, outside {workspace}"
print("rl_engine._C:", loaded)
PY
- name: Check pinned providers in their own process
# The FP32 router needs IEEE Triton dots from process start.
env:
TRITON_F32_DEFAULT: ieee
run: |
python3 -m pytest -q tests/models/qwen3_next/check_qwen3_next_norm_providers.py tests/models/qwen3_next/check_gdn_recurrent_golden.py \
tests/models/qwen3_next/check_qwen3_next_forward.py tests/models/qwen3_next/check_qwen3_next_attention.py \
tests/models/qwen3_next/check_qwen3_next_gdn_bridge.py tests/models/qwen3_next/check_qwen3_next_gdn_sequence.py \
tests/models/qwen3_next/check_qwen3_next_conv_bridge.py tests/models/qwen3_next/check_qwen3_next_core_matrix.py \
tests/models/qwen3_next/check_qwen3_next_shared_core.py
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