From 7dcd6a1ffacff1819d96475282abd509084e3366 Mon Sep 17 00:00:00 2001 From: RichApple123 <134492660+RichApple123@users.noreply.github.com> Date: Sat, 10 Oct 2026 18:40:22 +0800 Subject: [PATCH] feat(nemotron): add deterministic SM90 router dispatch Signed-off-by: RichApple123 <134492660+RichApple123@users.noreply.github.com> --- .github/workflows/ci.yml | 10 +- .../benchmark_nemotron_router_training.py | 307 ++++++++++ docs/.nav.yml | 1 + docs/operators/README.md | 1 + docs/operators/nemotron-router.md | 168 ++++++ .../kernels/ops/pytorch/nemotron_router.py | 103 ++++ .../kernels/ops/triton/nemotron_router.py | 524 ++++++++++++++++++ rl_engine/kernels/registry.py | 6 + tests/nemotron/test_nemotron_router_cuda.py | 427 ++++++++++++++ .../test_nemotron_router_reference.py | 129 +++++ 10 files changed, 1674 insertions(+), 2 deletions(-) create mode 100644 benchmarks/benchmark_nemotron_router_training.py create mode 100644 docs/operators/nemotron-router.md create mode 100644 rl_engine/kernels/ops/pytorch/nemotron_router.py create mode 100644 rl_engine/kernels/ops/triton/nemotron_router.py create mode 100644 tests/nemotron/test_nemotron_router_cuda.py create mode 100644 tests/nemotron/test_nemotron_router_reference.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 9c9b81539..f0aab5e2f 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-nemotron ] pull_request: - branches: [ main, test ] + branches: [ main, test, test-nemotron ] permissions: contents: read @@ -77,6 +77,12 @@ jobs: tests/test_tolerance_contract.py \ tests/test_kernel_registry.py + - name: Run Nemotron Router Contract Tests (CPU-safe) + run: | + PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest -q \ + tests/nemotron/test_nemotron_router_reference.py \ + tests/nemotron/test_nemotron_router_cuda.py + - name: Run Attention Ground-Truth Tests (CPU-safe) run: | python -m pytest tests/test_attention.py -v -k "not large and not gpu" diff --git a/benchmarks/benchmark_nemotron_router_training.py b/benchmarks/benchmark_nemotron_router_training.py new file mode 100644 index 000000000..f871b2455 --- /dev/null +++ b/benchmarks/benchmark_nemotron_router_training.py @@ -0,0 +1,307 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors +"""Operator-only cuBLAS + Transformer Engine 2.19 training baseline. + +Native probabilities retain the dense [T,E] mask ABI; candidate retains [T,K]. +Conversion for correctness and cotangent alignment is OUTSIDE timed work. +No expert GEMM, combine, auxiliary loss, or distributed communication included. +""" + +import argparse +import hashlib +import importlib.metadata +import json +import statistics +from pathlib import Path + +import torch +import transformer_engine.pytorch as te +from transformer_engine.pytorch.router import fused_topk_with_score_function + +import rl_engine.kernels.ops.triton.nemotron_router as core + +NATIVE_MAP_TYPE = "mask" + + +def routing_mask(route, n): + if route.dtype == torch.bool: + return route + return torch.zeros((n, 128), device=route.device, dtype=torch.bool).scatter_( + 1, route.long(), True + ) + + +def native(x, w, bias, full=True): + logits = torch.nn.functional.linear(x.float(), w) + indices = ( + torch.empty((x.shape[0], 6), device=x.device, dtype=torch.int32) + if NATIVE_MAP_TYPE == "index" + else None + ) + probs, mask = fused_topk_with_score_function( + logits, 6, False, None, None, 2.5, "sigmoid", bias, topk_indices=indices + ) + packed = ( + te.moe_permute(x, mask, x.shape[0] * 6, max_token_num=x.shape[0], map_type=NATIVE_MAP_TYPE)[ + 0 + ] + if full + else None + ) + return probs, mask, packed + + +def candidate(x, w, bias, full=True): + return core.nemotron_router_cuda(x, w, bias) if full else core._ProjectRoute.apply(x, w, bias) + + +def stats_error(a, b, atol, rtol): + diff = (a.double() - b.double()).abs() + return { + "allclose": bool(torch.allclose(a.double(), b.double(), atol=atol, rtol=rtol)), + "max_abs": diff.max().item(), + "relative_l2": (diff.norm() / b.double().norm().clamp_min(1e-30)).item(), + "atol": atol, + "rtol": rtol, + } + + +def qualify(x, w, bias, dense_grad, packed_grad, full): + a, b = native(x, w, bias, full), candidate(x, w, bias, full) + ids = b[0].long() + native_mask = routing_mask(a[1], x.shape[0]) + native_ids = native_mask.nonzero(as_tuple=True)[1].reshape(x.shape[0], 6) + same = (ids == native_ids).all(1) + report = {"same_expert_rows": int(same.sum()), "rows": x.shape[0]} + ga = dense_grad * native_mask + gb = dense_grad.gather(1, ids) + if full: + da = torch.autograd.grad((a[0], a[2]), (x, w), (ga, packed_grad)) + db = torch.autograd.grad((b[1], b[4]), (x, w), (gb, packed_grad)) + else: + da = torch.autograd.grad(a[0], (x, w), ga) + db = torch.autograd.grad(b[1], (x, w), gb) + # Cross-provider differences are meaningful only for an identical branch. + # Every provider is still qualified against FP64 on its OWN selected branch. + report["gradients_compared"] = bool(same.all()) + if bool(same.all()): + report["weights"] = stats_error(b[1], a[0].gather(1, ids), 2e-6, 2e-5) + report["dx"] = stats_error( + db[0], + da[0], + (0.0625 if full else 2e-4) if x.dtype == torch.bfloat16 else 2e-5, + 0.02 if x.dtype == torch.bfloat16 else 2e-4, + ) + report["dw"] = stats_error(db[1], da[1], 2e-4, 3e-4) + for name, selected, actual, actual_prob, payload in ( + ("native", native_ids, da, a[0].gather(1, native_ids), a[2]), + ("candidate", ids, db, b[1], b[4] if full else None), + ): + xx = x.detach().double().requires_grad_() + ww = w.detach().double().requires_grad_() + scores = torch.nn.functional.linear(xx, ww).sigmoid().gather(1, selected) + probs = scores / (scores.sum(1, keepdim=True) + 1e-20) * 2.5 + targets, cots = (probs,), (dense_grad.gather(1, selected).double(),) + if full: + perm = torch.argsort(selected.flatten(), stable=True) + targets += (xx[perm // 6],) + cots += (packed_grad.double(),) + report[name + "_payload_equal"] = bool(torch.equal(payload, x[perm // 6])) + ref = torch.autograd.grad(targets, (xx, ww), cots) + report[name + "_fp64_weights"] = stats_error(actual_prob, probs, 2e-6, 2e-5) + report[name + "_fp64_dx"] = stats_error( + actual[0], + ref[0], + (0.0625 if full else 2e-4) if x.dtype == torch.bfloat16 else 2e-5, + 0.02 if x.dtype == torch.bfloat16 else 2e-4, + ) + report[name + "_fp64_dw"] = stats_error(actual[1], ref[1], 2e-4, 3e-4) + assert all(report["candidate_fp64_" + key]["allclose"] for key in ("weights", "dx", "dw")) + if full: + assert report["candidate_payload_equal"] and report["native_payload_equal"] + return report + + +def capture(fn): + for _ in range(5): + fn() + torch.cuda.synchronize() + before = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + once = fn() + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - before + del once + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outputs = fn() + for _ in range(5): + graph.replay() + torch.cuda.synchronize() + return graph, outputs, peak + + +def paired_measure(calls, reverse, mode): + graphs = {} + if mode == "graph": + graphs = {name: capture(fn) for name, fn in calls.items()} + else: + for name, fn in calls.items(): + for _ in range(5): + fn() + torch.cuda.synchronize() + before = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + _once = fn() + assert _once is not None + torch.cuda.synchronize() + graphs[name] = (None, None, torch.cuda.max_memory_allocated() - before) + del _once + samples = {name: [] for name in calls} + order = list(calls) + if reverse: + order.reverse() + for i in range(20): + for name in order if i % 2 == 0 else order[::-1]: + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(5): + if mode == "graph": + graphs[name][0].replay() + else: + calls[name]() + end.record() + end.synchronize() + samples[name].append(start.elapsed_time(end) / 5) + return { + name: { + "median_ms": statistics.median(v), + "samples_ms": v, + "peak_increment_bytes": graphs[name][2], + } + for name, v in samples.items() + } + + +def main(): + global NATIVE_MAP_TYPE + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument( + "--tokens", type=int, nargs="+", default=[1, 16, 128, 1024, 4096, 8192, 32768] + ) + parser.add_argument("--dtype", choices=["bf16", "fp32"], default="bf16") + parser.add_argument("--distributions", nargs="+", default=["random", "concentrated"]) + parser.add_argument("--reverse", action="store_true") + parser.add_argument("--map-type", choices=["mask", "index"], default="mask") + parser.add_argument("--mode", choices=["graph", "eager"], default="graph") + args = parser.parse_args() + NATIVE_MAP_TYPE = args.map_type + if args.map_type == "index" and args.mode == "graph": + parser.error("TE 2.19 index radix sort uses the default stream; select --mode eager") + if args.output.exists(): + parser.error("refusing to overwrite results") + if importlib.metadata.version("transformer_engine") != "2.19.0": + parser.error("requires pinned Transformer Engine 2.19.0") + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.manual_seed(434) + result = { + "torch": torch.__version__, + "te": "2.19.0", + "gpu": torch.cuda.get_device_name(), + "source_sha256": hashlib.sha256(Path(core.__file__).read_bytes()).hexdigest(), + "benchmark_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), + "method": "20 alternating paired samples, 5 executions per sample; mode=" + + args.mode + + "; native dense ABI; no canonicalization timed; eager includes host launch gaps", + "scope": ( + "cuBLAS FP32 + TE fused sigmoid/top6 + TE " + + args.map_type + + " permutation and autograd; no expert/combine/aux loss/EP" + ), + "arguments": {**vars(args), "output": args.output.name}, + "rows": [], + } + dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float32 + # TE 2.19 index-map radix sort uses the default CUDA stream and is not + # safe to compare under a side-stream CUDA graph. Eager mode uses the + # default stream for BOTH providers, preserving upstream implementation. + execution_stream = torch.cuda.default_stream() if args.mode == "eager" else torch.cuda.Stream() + with torch.cuda.stream(execution_stream): + for distribution in args.distributions: + for n in args.tokens: + x = torch.randn(n, 2688, device="cuda", dtype=dtype).requires_grad_() + w = (torch.randn(128, 2688, device="cuda") * 0.02).requires_grad_() + bias = torch.randn(128, device="cuda") * 0.01 + if distribution == "concentrated": + bias.zero_() + bias[:6] = 4 + dense_grad = torch.randn(n, 128, device="cuda") + packed_grad = torch.randn(n * 6, 2688, device="cuda", dtype=dtype) + for full in (False, True): + qualification = qualify(x, w, bias, dense_grad, packed_grad, full) + out_a, out_b = native(x, w, bias, full), candidate(x, w, bias, full) + grads_a = ( + (dense_grad * routing_mask(out_a[1], n), packed_grad) + if full + else (dense_grad * routing_mask(out_a[1], n),) + ) + grads_b = ( + (dense_grad.gather(1, out_b[0].long()), packed_grad) + if full + else (dense_grad.gather(1, out_b[0].long()),) + ) + del out_a, out_b + + def fa(): + o = native(x, w, bias, full) + return (o[0], o[2]) if full else (o[0],) + + def fb(): + o = candidate(x, w, bias, full) + return (o[1], o[4]) if full else (o[1],) + + for phase in ("forward", "forward_backward"): + calls = ( + {"native": fa, "candidate": fb} + if phase == "forward" + else { + "native": lambda: torch.autograd.grad(fa(), (x, w), grads_a), + "candidate": lambda: torch.autograd.grad(fb(), (x, w), grads_b), + } + ) + measured = paired_measure(calls, args.reverse, args.mode) + row = { + "tokens": n, + "dtype": args.dtype, + "distribution": distribution, + "scope": "full" if full else "route", + "phase": phase, + "qualification": qualification, + **measured, + } + row["speedup"] = ( + measured["native"]["median_ms"] / measured["candidate"]["median_ms"] + ) + result["rows"].append(row) + args.output.write_text(json.dumps(result, indent=2)) + print( + n, + args.dtype, + distribution, + row["scope"], + phase, + "native", + round(measured["native"]["median_ms"], 6), + "candidate", + round(measured["candidate"]["median_ms"], 6), + "speedup", + round(row["speedup"], 4), + flush=True, + ) + print("TRAINING_BASELINE_DONE", flush=True) + + +if __name__ == "__main__": + main() diff --git a/docs/.nav.yml b/docs/.nav.yml index 7d71230a3..e86bf54c9 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -23,6 +23,7 @@ nav: - operators/batch-invariant-logp.md - operators/linear-logp-tp-test.md - operators/grpo-loss.md + - operators/nemotron-router.md - operators/lm_head.md - operators/ratio-kl.md - operators/pack-and-pad.md diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..5f9b9fa98 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -25,6 +25,7 @@ Every operator page should include: - [Batch-Invariant LogP](batch-invariant-logp.md) - [Fused Linear LogP TP Test Runbook](linear-logp-tp-test.md) - [GRPO Loss](grpo-loss.md) +- [Nemotron Nano Router and Dispatch](nemotron-router.md) - [RoPE](rope.md) - [LM Head](lm_head.md) - [Policy Ratio + KL Penalty](ratio-kl.md) diff --git a/docs/operators/nemotron-router.md b/docs/operators/nemotron-router.md new file mode 100644 index 000000000..1128515ed --- /dev/null +++ b/docs/operators/nemotron-router.md @@ -0,0 +1,168 @@ +# Nemotron Nano router and dispatch + +The SM90 CUDA provider implements the single-device `moe_router_dispatch` row +of [RFC #434](https://github.com/RL-Align/RL-Kernel/issues/434): projection, +sigmoid, corrected top-six selection, normalization, token packing and backward. +The ordering/arithmetic contract below is provisional pending maintainer review. +Expert MLP, weighted combine, shared +expert execution and EP communication are outside this operator. + +The reference checkpoint is `nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16`, +revision `bf77c31`, gate forward in `modeling_nemotron_h.py`. With one group, +all 128 experts remain eligible. This proposal is `nemotron-router-sm90-v3`. +The maintainer-owned architecture fingerprint and router section 4.4 remain +pending for canonical acceptance. Implementation was authorized in +[the issue discussion](https://github.com/RL-Align/RL-Kernel/issues/434#issuecomment-5977907553). +Consumers must review the proposed ordering before assuming this ABI. + +## Interface + +```python +from rl_engine.kernels.registry import KernelRegistry + +op = KernelRegistry().get_op("nemotron_router_dispatch", device=x.device) +ids, weights, permutation, offsets, packed = op(x, w, bias) +``` + +All inputs must be contiguous, finite and on the same NVIDIA SM90 device. +Autocast is not used inside the operator; unsupported geometry/dtypes fail. + +| Tensor | Shape | Dtype | Meaning | +| --- | --- | --- | --- | +| `x` | `[T,2688]` | BF16 or FP32 | Token rows | +| `w` | `[128,2688]` | FP32 | Router weight | +| `bias` | `[128]` | FP32 | Selection-only correction | +| `ids` | `[T,6]` | int64 | Selected experts in ascending ID order | +| `weights` | `[T,6]` | FP32 | Normalized original sigmoid scores times 2.5 | +| `permutation` | `[6*T]` | int64 | Packed position to original flat route slot | +| `offsets` | `[129]` | int64 | Exclusive expert prefix counts | +| `packed` | `[6*T,2688]` | Same as `x` | Unweighted token copies | + +`0 <= T <= 65536`; larger counts fail before allocation to bound fixed-width +index products. There is no padding, dropping or capacity truncation. Packing +is expert-major then token-major; each permutation value is `token * 6 + slot`. +Offsets delimit expert segments. Integer block prefix sums give each route a +unique scatter destination without atomics. Absolute positions change with the +batch, so payload invariance is checked after undoing the permutation. + +The registry has no CPU/ROCm or non-strict fallback. The hot path checks metadata +only; callers must ensure finite inputs and intermediates. The separate +`route_scores_cuda` debugging interface synchronously checks score values. + +## Fixed arithmetic + +Projection transposes the current weight to contiguous `[2688,128]` on every +call, without a weight cache. Fixed 32-by-64 tiles compute 21 independent +128-element K segments, each visiting four ordered 32-wide IEEE FP32 dot blocks. +A zero-padded 32-leaf FP32 tree merges them. Two warps and three stages are +fixed across token counts. TF32, autotuning, atomic sums and FP fusion are disabled. +Merge, FP32 sigmoid `1 / (1 + exp(-logit))` and routing are fused. + +Ranking uses `score + bias`; exact ties choose the lower expert ID. Six argmax +reductions select distinct experts, then an eight-lane sort orders the six IDs. +Normalization sums the original sigmoid scores in ascending selected-ID order, +adds `1e-20`, divides and multiplies by 2.5. Bias has no gradient. + +Routing backward fuses the sigmoid derivative as `(dscore * (1 - score)) * score`. +dX and dW use fixed 64-by-64 tiles with ordered 16-wide IEEE FP32 dot blocks and +four warps. dX has one 128-element segment. dW uses 512-token segments and a +zero-padded next-power-of-two sum tree. A single segment skips the merge; short +single segments run `ceil(K/16)` blocks, while multiple segments run 32 blocks +with masked tails. Empty reductions produce zero. + +Payload backward gathers six gradients in selected-slot order in FP32. +Routing and payload dX branches each round to the input dtype before addition. +dW reduces the complete token set in its supplied order; adding independently +computed microbatch gradients is not a bitwise full-batch equivalence claim. +Top-k membership is discrete: gradients flow through the selected scores and +copied payloads, not indices or bias. Higher-order gradients are unqualified. + +The target branch's `TritonDetGemmOp.forward_fp32` uses a different projection +reduction order and its backward rejects FP32 router weights through the BF16 +tree path. Reusing it would require an agreed arithmetic contract and a qualified +FP32 backward. The specialized implementation is restricted to the geometry above. + +## Qualification boundary + +On the same architecture, compiler and software stack, an identical token and +upstream gradient must preserve IDs, weights, payload and dX bits across batch +sizes and positions. Payloads are compared after undoing dispatch. dW is +repeatable for an identical ordered complete token set. Microbatch dW addition, +cross-compiler equivalence, ROCm and TP/CP/EP equivalence are unqualified. +Native GEMM/TE/HF providers need not produce identical bits under this proposed +fixed arithmetic. Consumers using other route conventions need an explicit adapter; +the registry entry alone does not provide framework integration. + +## Tests + +```bash +python -m pytest tests/test_kernel_registry.py tests/nemotron/ -q +``` + +Tests cover an independent scalar FP64 reference, reference finite-difference +gradcheck away from selection boundaries, CUDA accuracy, ties, empty inputs, +unused-output/frozen-input gradients, raw-bit row/dX invariance, fixed-set dW +repeatability, tails, cancellation, the token-count bound, CUDA Graph weight +replacement and outer autocast. The consumer test applies independent nonlinear +experts and weighted combine to expose packing, weighting and backward errors; +it does not implement or benchmark the Nemotron expert MLP. + +Test tolerances are empirical workload thresholds, not universal error bounds +or maintainer-approved thresholds. Near a top-k boundary FP32 rounding can +change selection; each provider's selected branch is qualified against FP64. + +Current source SHA256: +`9ec652f622711d303b112787f3405751ddc7a5ed4f0281f74fb172f970605a46`. +On one H20 with Torch 2.9.1+cu130 / Triton 3.5.1, 106 reference, registry, +CUDA and consumer tests pass. All five outputs and dX/dW match the previous +implementation bitwise in 36 direct cases through 65536 tokens. Compute +Sanitizer 2025.3.1.0 passes all four tools on full/tail/capacity workloads. +These records cover selected operator sources/tests, not full-repository GPU CI. + +## Performance reproduction + +The primary production baseline is FP32 cuBLAS projection plus unchanged +Transformer Engine 2.19 fused sigmoid/top-six and index permutation/autograd. +It preserves native formats, with no timed canonicalization. Cotangents are +aligned outside timing. Expert GEMM, combine, auxiliary loss, optimizer and +communication are excluded from both providers. + +```bash +# Install optional TE only in an isolated CUDA environment. +pip install 'transformer_engine[pytorch]==2.19.0' +PYTHONPATH=. python benchmarks/benchmark_nemotron_router_training.py --map-type index --mode eager --dtype bf16 --output /tmp/router-te-bf16.json +# Repeat with --dtype fp32 and --reverse, using new output paths. +``` + +The TE index implementation does not forward the current stream to its CUB +sort, so the comparison uses default-stream eager timing for both providers. +It includes host-launch gaps. Do not combine CUDA Graph latencies from other +experiments with these eager measurements into a speedup. +TE used its official cu13 wheel and PyTorch binding with optional EP bindings +disabled. The measured process excluded an unused incompatible optional FA4 +import; router/permutation code was unchanged. An isolated environment without +FA4 avoids that import conflict. + +The current evidence snapshot, supplied separately with the PR as +`nemotron-router-current-h20.json`, contains all 128 full-operator index/eager measurements: eight token counts +(1/16/128/513/1024/4096/8192/32768), both dtypes, random/concentrated routing, +forward/forward+backward and both initial provider orders. Each measurement +retains 20 alternating paired samples of five executions, medians, incremental +peak allocation and FP64 qualification. It binds implementation/benchmark hashes +and preserves regressions. Broader local qualification covered 512 measurements; +only the current conservative full-operator matrix is included here. + +Full forward+backward on random inputs, in milliseconds. Latencies use the +normal initial order; speedup ranges cover both initial orders: + +| Dtype | Tokens | Native ms | Strict ms | Native / strict | +| --- | --- | --- | --- | --- | +| bf16 | 8192 | 1.235402 | 1.141997 | 1.082-1.083x | +| bf16 | 32768 | 4.603255 | 4.217085 | 1.090-1.092x | +| fp32 | 8192 | 1.310125 | 1.293034 | 1.013-1.014x | +| fp32 | 32768 | 4.878426 | 4.857485 | 1.004-1.005x | + +BF16 large-token training has an 8-9% speedup ratio; FP32's small difference is +near parity and does not establish robust acceleration. Small-token regressions +remain in the matrix. No universal speedup, model-quality improvement, +distributed equivalence or full-model throughput result is claimed. diff --git a/rl_engine/kernels/ops/pytorch/nemotron_router.py b/rl_engine/kernels/ops/pytorch/nemotron_router.py new file mode 100644 index 000000000..f5a5d65f1 --- /dev/null +++ b/rl_engine/kernels/ops/pytorch/nemotron_router.py @@ -0,0 +1,103 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors +"""Nemotron Nano routing specification/reference, not a strict GPU provider. + +The ordering/dispatch ABI is a proposal for RFC #434. GEMM and autograd here +are numerical baselines only: CUDA matmul and gather backward are not claimed +batch invariant. The independent scalar oracle belongs in tests. +""" + +from dataclasses import dataclass + +import torch + + +@dataclass +class RouterResult: + expert_ids: torch.Tensor # [T, K], ascending expert id within each row + weights: torch.Tensor # [T, K], aligned with expert_ids + permutation: torch.Tensor # [T*K], packed position -> original flat route slot + expert_offsets: torch.Tensor # [E+1], exclusive prefix counts + packed_tokens: torch.Tensor # [T*K, H], unweighted + + +def route_scores(scores, correction_bias, top_k=6, scaling=2.5): + """Select by corrected scores; weight using original scores. + + Proposed tie policy: lower expert id wins. Selected experts are then + sorted by id; normalization sums in that order, then adds 1e-20. + Top-k selection is discrete; bias has no differentiable output path. + """ + if scores.ndim != 2 or not scores.is_floating_point(): + raise ValueError("scores must be a floating [tokens, experts] tensor") + if correction_bias.shape != (scores.shape[1],): + raise ValueError("correction_bias must have shape [experts]") + if correction_bias.device != scores.device or correction_bias.dtype != scores.dtype: + raise ValueError("scores and correction_bias must share dtype and device") + if not 1 <= top_k <= scores.shape[1]: + raise ValueError("top_k must be between 1 and the number of experts") + if not torch.isfinite(scores).all() or not torch.isfinite(correction_bias).all(): + raise ValueError("non-finite routing inputs are unsupported") + if (scores < 0).any() or (scores > 1).any(): + raise ValueError("scores must lie in [0, 1]") + with torch.no_grad(): + choice = scores + correction_bias + if not torch.isfinite(choice).all(): + raise ValueError("corrected score overflow") + ids = torch.argsort(choice, dim=-1, descending=True, stable=True)[:, :top_k] + ids = ids.sort(dim=-1).values + selected = scores.gather(1, ids) + denominator = torch.zeros_like(selected[:, :1]) + for slot in range(top_k): + denominator = denominator + selected[:, slot : slot + 1] + weights = (selected / (denominator + 1e-20)) * scaling + return ids, weights + + +def dispatch_tokens(hidden_states, expert_ids, num_experts): + """Proposed ABI: expert-major, stable token-major unweighted dispatch. + + Permutation stores flattened (token, route-slot) indices, retaining enough + information to invert dispatch. No padding, capacity truncation or dropping. + """ + if hidden_states.ndim != 2 or expert_ids.ndim != 2: + raise ValueError("hidden_states and expert_ids must be rank two") + if expert_ids.shape[0] != hidden_states.shape[0]: + raise ValueError("token counts must agree") + if expert_ids.dtype != torch.int64 or expert_ids.device != hidden_states.device: + raise ValueError("expert_ids must be int64 on the input device") + if num_experts < 1 or expert_ids.shape[1] < 1: + raise ValueError("positive expert and route counts are required") + if ((expert_ids < 0) | (expert_ids >= num_experts)).any(): + raise ValueError("expert id out of range") + if expert_ids.shape[1] > 1 and (expert_ids[:, 1:] <= expert_ids[:, :-1]).any(): + raise ValueError("expert ids must be unique and ascending per token") + flat = expert_ids.reshape(-1) + permutation = torch.argsort(flat, stable=True) + counts = torch.bincount(flat, minlength=num_experts) + offsets = torch.cat((counts.new_zeros(1), counts.cumsum(0))) + packed = hidden_states.index_select(0, permutation // expert_ids.shape[1]) + return permutation, offsets, packed + + +def router_reference(hidden_states, router_weight, correction_bias, top_k=6, scaling=2.5): + """Differentiable baseline. FP64 inputs retain FP64 for gradient checks. + + Otherwise projection, sigmoid and routing weights use FP32 as in the + pinned checkpoint. Packed payload retains the input dtype. + """ + if hidden_states.ndim != 2 or router_weight.ndim != 2: + raise ValueError("expected [T,H] input and [E,H] router weights") + if hidden_states.shape[1] != router_weight.shape[1]: + raise ValueError("hidden widths must agree") + if not hidden_states.is_floating_point() or not router_weight.is_floating_point(): + raise ValueError("input and weight must be floating tensors") + dtype = ( + torch.float64 + if hidden_states.dtype == router_weight.dtype == torch.float64 + else torch.float32 + ) + logits = hidden_states.to(dtype) @ router_weight.to(dtype).T + ids, weights = route_scores(logits.sigmoid(), correction_bias.to(dtype), top_k, scaling) + permutation, offsets, packed = dispatch_tokens(hidden_states, ids, router_weight.shape[0]) + return RouterResult(ids, weights, permutation, offsets, packed) diff --git a/rl_engine/kernels/ops/triton/nemotron_router.py b/rl_engine/kernels/ops/triton/nemotron_router.py new file mode 100644 index 000000000..da936f6d9 --- /dev/null +++ b/rl_engine/kernels/ops/triton/nemotron_router.py @@ -0,0 +1,524 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors +"""Nemotron Nano single-device router with provisional RFC #434 ordering. + +SM90 CUDA only. Forward and per-token gradients preserve a fixed FP32 tree; +weight gradients use the complete logical token order. No atomic reductions. +""" + +import torch +import triton +import triton.language as tl + +CONTRACT_VERSION = "nemotron-router-sm90-v3" +# Bound every 32-bit index product, including the six-copy packed payload. +MAX_TOKENS = 65536 + + +@triton.jit +def _route_from_scores(s, B, IDS, W, D, t): + e = tl.arange(0, 128) + lane = tl.arange(0, 8) + choices = s + tl.load(B + e) + choices = tl.where(choices == choices, choices, -float("inf")) + ids = tl.full((8,), 128, tl.int32) + for slot in range(6): + _, winner = tl.max(choices, 0, return_indices=True, return_indices_tie_break_left=True) + ids = tl.where(lane == slot, winner, ids) + choices = tl.where(e == winner, -float("inf"), choices) + # Only eight lanes need sorting/gathering after six deterministic argmaxes. + ids = tl.sort(ids, descending=False) + values = tl.gather(s, tl.minimum(ids, 127), 0) + denom = tl.full((), 0.0, tl.float32) + for slot in range(6): + denom = denom + tl.sum(tl.where(lane == slot, values, 0.0), 0) + denom = denom + 1.0e-20 + tl.store(D + t, denom) + tl.store(IDS + t * 6 + lane, ids, lane < 6) + tl.store(W + t * 6 + lane, (values / denom) * 2.5, lane < 6) + + +@triton.jit +def _route_fwd(S, B, IDS, W, D): + t = tl.program_id(0) + e = tl.arange(0, 128) + _route_from_scores(tl.load(S + t * 128 + e), B, IDS, W, D, t) + + +@triton.jit +def _merge_and_route(P, B, SCORES, IDS, W, D, T: tl.constexpr): + t = tl.program_id(0) + e = tl.arange(0, 128) + segment = tl.arange(0, 32) + partials = tl.load( + P + segment[:, None] * T * 128 + t * 128 + e[None, :], segment[:, None] < 21, 0 + ) + logits = tl.sum(partials, 0) + scores = 1.0 / (1.0 + tl.exp(-logits)) + tl.store(SCORES + t * 128 + e, scores) + _route_from_scores(scores, B, IDS, W, D, t) + + +@triton.jit +def _route_bwd(S, IDS, D, G, DS, SIGMOID: tl.constexpr = False): + t = tl.program_id(0) + e = tl.arange(0, 128) + denom = tl.load(D + t) + dot = tl.full((), 0.0, tl.float32) + for slot in range(6): + idx = tl.load(IDS + t * 6 + slot) + score = tl.load(S + t * 128 + idx) + grad = tl.load(G + t * 6 + slot) + dot = dot + grad * score + result = tl.full((128,), 0.0, tl.float32) + for slot in range(6): + idx = tl.load(IDS + t * 6 + slot) + grad = tl.load(G + t * 6 + slot) + ds = (grad / denom - (dot / denom) / denom) * 2.5 + result = tl.where(e == idx, ds, result) + if SIGMOID: + scores = tl.load(S + t * 128 + e) + result = (result * (1.0 - scores)) * scores + tl.store(DS + t * 128 + e, result) + + +class _RouteScores(torch.autograd.Function): + @staticmethod + def forward(ctx, scores, bias): + n = scores.shape[0] + ids = torch.empty((n, 6), device=scores.device, dtype=torch.int64) + weights = torch.empty((n, 6), device=scores.device, dtype=torch.float32) + denom = torch.empty(n, device=scores.device, dtype=torch.float32) + if n: + _route_fwd[(n,)](scores, bias, ids, weights, denom, num_warps=4, enable_fp_fusion=False) + ctx.save_for_backward(scores, ids, denom) + ctx.mark_non_differentiable(ids) + return ids, weights + + @staticmethod + def backward(ctx, grad_ids, grad_weights): + scores, ids, denom = ctx.saved_tensors + ds = torch.empty_like(scores) + if scores.shape[0]: + _route_bwd[(scores.shape[0],)]( + scores, + ids, + denom, + grad_weights.contiguous(), + ds, + num_warps=4, + enable_fp_fusion=False, + ) + return ds, None + + +def route_scores_cuda(scores, correction_bias): + """Checked experimental interface. Validation syncs; benchmark separately. + + Non-finite inputs fail closed, including overflow of corrected scores. + This primitive implements only the fixed Nano 128-expert/top-6 contract. + """ + if not scores.is_cuda or scores.dtype != torch.float32: + raise ValueError("FP32 CUDA scores required") + if scores.ndim != 2 or scores.shape[1] != 128 or not scores.is_contiguous(): + raise ValueError("contiguous [T,128] scores required") + if ( + correction_bias.shape != (128,) + or correction_bias.dtype != torch.float32 + or correction_bias.device != scores.device + or not correction_bias.is_contiguous() + ): + raise ValueError("contiguous FP32 bias[128] on input device required") + if ( + not torch.isfinite(scores).all() + or not torch.isfinite(correction_bias).all() + or not torch.isfinite(scores + correction_bias).all() + or (scores < 0).any() + or (scores > 1).any() + ): + raise ValueError("finite sigmoid scores and non-overflowing corrected scores required") + return _RouteScores.apply(scores, correction_bias) + + +@triton.jit +def _backward_partials( + A, B, P, M: tl.constexpr, N: tl.constexpr, K: tl.constexpr, CH: tl.constexpr, BN: tl.constexpr +): + # Fixed output tiles and ordered IEEE FP32 products within each K segment. + m = tl.program_id(0) * 64 + tl.arange(0, 64) + n = tl.program_id(1) * BN + tl.arange(0, BN) + k = tl.arange(0, 16) + acc = tl.full((64, BN), 0, tl.float32) + for block in range(tl.cdiv(tl.minimum(K, CH), 16)): + kk = tl.program_id(2) * CH + block * 16 + k + a = tl.load(A + m[:, None] * K + kk[None, :], (m[:, None] < M) & (kk[None, :] < K), 0).to( + tl.float32 + ) + b = tl.load(B + kk[:, None] * N + n[None, :], (kk[:, None] < K) & (n[None, :] < N), 0).to( + tl.float32 + ) + acc = tl.dot(a, b, acc, input_precision="ieee") + tl.store( + P + tl.program_id(2) * M * N + m[:, None] * N + n[None, :], + acc, + (m[:, None] < M) & (n[None, :] < N), + ) + + +@triton.jit +def _backward_merge(P, Y, M: tl.constexpr, N: tl.constexpr, S: tl.constexpr, R: tl.constexpr): + i = tl.program_id(0) * 128 + tl.arange(0, 128) + s = tl.arange(0, R) + values = tl.load( + P + s[:, None] * M * N + i[None, :], (s[:, None] < S) & (i[None, :] < M * N), 0 + ) + tl.store(Y + i, tl.sum(values, 0), i < M * N) + + +def _mm(a, b, chunk, out_dtype=torch.float32): + a, b = a.contiguous(), b.contiguous() + m, k = a.shape + n = b.shape[1] + out = torch.empty((m, n), device=a.device, dtype=out_dtype) + if not m or not n: + return out + if not k: + return out.zero_() + segments = triton.cdiv(k, chunk) + partials = ( + out + if segments == 1 + else torch.empty((segments, m, n), device=a.device, dtype=torch.float32) + ) + # Narrow dX tiles reduce per-thread accumulator pressure without changing K order. + bn = 64 + _backward_partials[(triton.cdiv(m, 64), triton.cdiv(n, bn), segments)]( + a, + b, + partials, + m, + n, + k, + chunk, + bn, + num_warps=4, + # A single stage reduces BF16 weight-gradient resource pressure. + # Keep the K blocks, output tiles and reduction tree unchanged. + num_stages=1 if b.dtype == torch.bfloat16 else 3, + enable_fp_fusion=False, + ) + if segments > 1: + _backward_merge[(triton.cdiv(m * n, 128),)]( + partials, + out, + m, + n, + segments, + triton.next_power_of_2(segments), + num_warps=4, + enable_fp_fusion=False, + ) + return out + + +@triton.jit +def _project_partials(X, WT, PARTIALS, T: tl.constexpr): + # Fixed 32x64 output tiles and 21 independent 128-element K segments. + # WT is contiguous [2688,128]; each segment visits four 32-wide blocks. + m = tl.program_id(0) * 32 + tl.arange(0, 32) + n = tl.program_id(1) * 64 + tl.arange(0, 64) + k = tl.arange(0, 32) + segment = tl.program_id(2) + acc = tl.full((32, 64), 0.0, tl.float32) + for block in range(4): + kk = segment * 128 + block * 32 + k + x = tl.load(X + m[:, None] * 2688 + kk[None, :], m[:, None] < T, 0) + w = tl.load(WT + kk[:, None] * 128 + n[None, :]) + acc = tl.dot(x.to(tl.float32), w, acc, input_precision="ieee") + tl.store( + PARTIALS + segment * T * 128 + m[:, None] * 128 + n[None, :], + acc, + m[:, None] < T, + ) + + +@triton.jit +def _project_merge(PARTIALS, OUT, T: tl.constexpr): + # A fixed 32-leaf FP32 tree: segment leaves 21..31 are exact zero. + n = tl.program_id(0) * 128 + tl.arange(0, 128) + segment = tl.arange(0, 32) + partials = tl.load( + PARTIALS + segment[:, None] * T * 128 + n[None, :], + segment[:, None] < 21, + 0, + ) + tl.store(OUT + n, tl.sum(partials, 0)) + + +class _Projection(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight): + ctx.save_for_backward(x, weight) + out = torch.empty((x.shape[0], 128), device=x.device, dtype=torch.float32) + if x.shape[0]: + transposed_weight = weight.t().contiguous() + partials = torch.empty((21, x.shape[0], 128), device=x.device, dtype=torch.float32) + _project_partials[(triton.cdiv(x.shape[0], 32), 2, 21)]( + x, + transposed_weight, + partials, + x.shape[0], + num_warps=2, + num_stages=3, + enable_fp_fusion=False, + ) + _project_merge[(x.shape[0],)]( + partials, out, x.shape[0], num_warps=4, enable_fp_fusion=False + ) + return out + + @staticmethod + def backward(ctx, grad): + x, weight = ctx.saved_tensors + dx = _mm(grad, weight, 128).to(x.dtype) if ctx.needs_input_grad[0] else None + dw = _mm(grad.t(), x, 512).to(weight.dtype) if ctx.needs_input_grad[1] else None + return dx, dw + + +def _project_route_forward(x, weight, bias): + """Shared executable forward; qualification reads the actual fused scores.""" + n = x.shape[0] + scores = torch.empty((n, 128), device=x.device, dtype=torch.float32) + ids = torch.empty((n, 6), device=x.device, dtype=torch.int64) + weights = torch.empty((n, 6), device=x.device, dtype=torch.float32) + denom = torch.empty(n, device=x.device, dtype=torch.float32) + if n: + wt = weight.t().contiguous() + partials = torch.empty((21, n, 128), device=x.device, dtype=torch.float32) + _project_partials[(triton.cdiv(n, 32), 2, 21)]( + x, wt, partials, n, num_warps=2, num_stages=3, enable_fp_fusion=False + ) + _merge_and_route[(n,)]( + partials, bias, scores, ids, weights, denom, n, num_warps=4, enable_fp_fusion=False + ) + return ids, weights, scores, denom + + +class _ProjectRoute(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight, bias): + ids, weights, scores, denom = _project_route_forward(x, weight, bias) + ctx.save_for_backward(x, weight, scores, ids, denom) + ctx.mark_non_differentiable(ids) + return ids, weights + + @staticmethod + def backward(ctx, grad_ids, grad_weights): + x, weight, scores, ids, denom = ctx.saved_tensors + grad = torch.empty_like(scores) + if x.shape[0]: + _route_bwd[(x.shape[0],)]( + scores, + ids, + denom, + grad_weights.contiguous(), + grad, + True, + num_warps=4, + enable_fp_fusion=False, + ) + dx = _mm(grad, weight, 128).to(x.dtype) if ctx.needs_input_grad[0] else None + dw = _mm(grad.t(), x, 512).to(weight.dtype) if ctx.needs_input_grad[1] else None + return dx, dw, None + + +@triton.jit +def _dispatch_counts(IDS, Counts, T: tl.constexpr, BLOCKS: tl.constexpr): + block = tl.program_id(0) + r = block * 256 + tl.arange(0, 256) + ids = tl.load(IDS + r, r < T * 6, 0).to(tl.int32) + # Invalid tail lanes must be excluded explicitly from the histogram. + counts = tl.histogram(ids, 128, mask=r < T * 6) + tl.store(Counts + tl.arange(0, 128) * BLOCKS + block, counts) + + +@triton.jit +def _integer_maximum(a, b): + return tl.maximum(a, b) + + +@triton.jit +def _dispatch_map(IDS, Prefix, Offsets, Perm, Inverse, T: tl.constexpr, BLOCKS: tl.constexpr): + block = tl.program_id(0) + lane = tl.arange(0, 256) + r = block * 256 + lane + ids = tl.load(IDS + r, r < T * 6, 128).to(tl.int32) + # Unique integer keys sort by expert, then original route position. The + # sentinel expert 128 sorts tail lanes last and never writes an output. + key = tl.sort(ids * 256 + lane, descending=False) + expert = key // 256 + original = block * 256 + key % 256 + previous = tl.gather(expert, tl.maximum(lane - 1, 0), 0) + starts = tl.where((lane == 0) | (expert != previous), lane, 0) + group_start = tl.associative_scan(starts, 0, _integer_maximum) + before = tl.load(Prefix + expert * BLOCKS + block - 1, (expert < 128) & (block > 0), 0) + # Prefix counts of earlier blocks plus local stable rank recover the + # global expert-major / route-major position without floating atomics. + dest = tl.load(Offsets + expert, expert < 128, 0) + before + lane - group_start + tl.store(Perm + dest, original, expert < 128) + tl.store(Inverse + original, dest, expert < 128) + + +@triton.jit +def _pack(X, Inverse, Y, H: tl.constexpr): + token = tl.program_id(0) + h = tl.program_id(1) * 1024 + tl.arange(0, 1024) + # Read a token once and reuse its exact payload for all six destinations. + # Inverse is a bijection, so CTAs never race on an output address. + value = tl.load(X + token * H + h, h < H, 0) + for slot in tl.static_range(6): + dest = tl.load(Inverse + token * 6 + slot) + tl.store(Y + dest * H + h, value, h < H) + + +@triton.jit +def _unpack_grad(G, Inverse, DX, H: tl.constexpr, Route=None, HAS_ROUTE: tl.constexpr = False): + token = tl.program_id(0) + h = tl.program_id(1) * 1024 + tl.arange(0, 1024) + total = tl.full((1024,), 0, tl.float32) + for slot in range(6): + r = tl.load(Inverse + token * 6 + slot) + total = total + tl.load(G + r * H + h, h < H, 0).to(tl.float32) + if HAS_ROUTE: + # Preserve separate branch rounding before the autograd-equivalent sum. + total = total.to(DX.dtype.element_ty).to(tl.float32) + total = total + tl.load(Route + token * H + h, h < H, 0).to(tl.float32) + tl.store(DX + token * H + h, total, h < H) + + +def _dispatch_forward(x, ids): + t, h = x.shape + r = t * 6 + blocks = triton.cdiv(r, 256) + perm = torch.empty(r, device=x.device, dtype=torch.int64) + inverse = torch.empty_like(perm) + offsets = torch.zeros(129, device=x.device, dtype=torch.int64) + packed = torch.empty((r, h), device=x.device, dtype=x.dtype) + if t: + counts = torch.empty((128, blocks), device=x.device, dtype=torch.int64) + _dispatch_counts[(blocks,)](ids, counts, t, blocks) + prefix = counts.cumsum(1) + offsets[1:] = prefix[:, -1].cumsum(0) + _dispatch_map[(blocks,)](ids, prefix, offsets, perm, inverse, t, blocks) + _pack[(t, triton.cdiv(h, 1024))](x, inverse, packed, h) + return perm, offsets, packed, inverse + + +class _Dispatch(torch.autograd.Function): + @staticmethod + def forward(ctx, x, ids): + t, h = x.shape + perm, offsets, packed, inverse = _dispatch_forward(x, ids) + ctx.save_for_backward(inverse) + ctx.shape = (t, h) + ctx.mark_non_differentiable(perm, offsets) + return perm, offsets, packed + + @staticmethod + def backward(ctx, gperm, goffsets, grad): + (inverse,) = ctx.saved_tensors + t, h = ctx.shape + dx = torch.empty((t, h), device=grad.device, dtype=grad.dtype) + if t: + _unpack_grad[(t, triton.cdiv(h, 1024))]( + grad.contiguous(), inverse, dx, h, enable_fp_fusion=False + ) + return dx, None + + +class _RouterDispatch(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight, bias): + ids, weights, scores, denom = _project_route_forward(x, weight, bias) + perm, offsets, packed, inverse = _dispatch_forward(x, ids) + ctx.save_for_backward(x, weight, scores, ids, denom, inverse) + ctx.mark_non_differentiable(ids, perm, offsets) + ctx.set_materialize_grads(False) + return ids, weights, perm, offsets, packed + + @staticmethod + def backward(ctx, grad_ids, grad_weights, grad_perm, grad_offsets, grad_packed): + x, weight, scores, ids, denom, inverse = ctx.saved_tensors + dx = dw = None + if grad_weights is not None: + grad = torch.empty_like(scores) + if x.shape[0]: + _route_bwd[(x.shape[0],)]( + scores, + ids, + denom, + grad_weights.contiguous(), + grad, + True, + num_warps=4, + enable_fp_fusion=False, + ) + if ctx.needs_input_grad[0]: + dx = _mm(grad, weight, 128, out_dtype=x.dtype) + if ctx.needs_input_grad[1]: + dw = _mm(grad.t(), x, 512) + if grad_packed is not None and ctx.needs_input_grad[0]: + combined = torch.empty_like(x) + if x.shape[0]: + _unpack_grad[(x.shape[0], triton.cdiv(x.shape[1], 1024))]( + grad_packed.contiguous(), + inverse, + combined, + x.shape[1], + Route=dx, + HAS_ROUTE=dx is not None, + enable_fp_fusion=False, + ) + dx = combined + return dx, dw, None + + +def nemotron_router_cuda(x, weight, correction_bias): + """Single-device Nano router with a provisional deterministic ordering ABI. + + Inputs must be finite. Shapes/dtypes are validated without device sync. + Weight gradients use increasing logical token order; externally accumulated + microbatch weight gradients are not promised bitwise equivalence. + """ + if not x.is_cuda or x.ndim != 2 or x.shape[1] != 2688: + raise ValueError("CUDA x[T,2688] required") + if x.shape[0] > MAX_TOKENS: + raise ValueError(f"at most {MAX_TOKENS} tokens supported by this provider") + if x.dtype not in (torch.bfloat16, torch.float32) or not x.is_contiguous(): + raise ValueError("contiguous BF16 or FP32 input required") + if ( + weight.shape != (128, 2688) + or weight.dtype != torch.float32 + or weight.device != x.device + or not weight.is_contiguous() + ): + raise ValueError("contiguous FP32 weight[128,2688] on input device required") + if ( + correction_bias.shape != (128,) + or correction_bias.dtype != torch.float32 + or correction_bias.device != x.device + or not correction_bias.is_contiguous() + ): + raise ValueError("contiguous FP32 bias[128] on input device required") + if torch.version.hip is not None or torch.cuda.get_device_capability(x.device) != ( + 9, + 0, + ): + raise ValueError("this provider is qualified only for NVIDIA SM90") + return _RouterDispatch.apply(x, weight, correction_bias) + + +class NemotronRouterOp: + """Registry entry point for the provisional single-device Nano router.""" + + __call__ = staticmethod(nemotron_router_cuda) diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 9eeb1c42b..e9036955b 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -90,6 +90,9 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): "rl_engine.kernels.ops.rocm.attention.flash_attn.StrictRocmAiterCKAttentionCore" ) + # Nemotron Nano single-device router; no reference fallback in strict dispatch. + TRITON_NEMOTRON_ROUTER = "rl_engine.kernels.ops.triton.nemotron_router.NemotronRouterOp" + # GRPO loss (group reward normalization + clipped surrogate + KL) TRITON_GRPO_LOSS = "rl_engine.kernels.ops.triton.loss.grpo_loss.TritonGRPOLossOp" PYTORCH_GRPO_LOSS = "rl_engine.kernels.ops.pytorch.loss.grpo_loss.NativeGRPOLossOp" @@ -535,6 +538,7 @@ def __init__(self): self._priority_map = { "cuda": { + "nemotron_router_dispatch": [OpBackend.TRITON_NEMOTRON_ROUTER], "logp": [ OpBackend.CUDA_FUSED_LOGP_GENERIC, OpBackend.FLASHINFER, @@ -1034,6 +1038,8 @@ def get_op(self, op_type: str, device: torch.device | str | None = None) -> Any: """Select the best legacy operator for the requested device.""" platform = self._platform_for_device(device) + if op_type == "nemotron_router_dispatch" and platform != "cuda": + raise RuntimeError("Nemotron router is qualified only for NVIDIA SM90 CUDA") candidates = self._priority_map.get(platform, {}).get(op_type, [OpBackend.PYTORCH_NATIVE]) for backend in candidates: diff --git a/tests/nemotron/test_nemotron_router_cuda.py b/tests/nemotron/test_nemotron_router_cuda.py new file mode 100644 index 000000000..68fb8b5e2 --- /dev/null +++ b/tests/nemotron/test_nemotron_router_cuda.py @@ -0,0 +1,427 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors +"""CUDA router accuracy, invariance, boundary and consumer qualification.""" + +import pytest +import torch + +pytestmark = [ + pytest.mark.cuda_only, + pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required"), +] + +_SM90 = pytest.mark.skipif( + not torch.cuda.is_available() + or torch.version.hip is not None + or torch.cuda.get_device_capability() != (9, 0), + reason="NVIDIA SM90 required", +) + + +@pytest.mark.parametrize("n", [0, 1, 7, 128, 1025]) +def test_scores_forward_backward(n): + pytest.importorskip("triton") + from rl_engine.kernels.ops.pytorch.nemotron_router import route_scores + from rl_engine.kernels.ops.triton.nemotron_router import route_scores_cuda + + torch.manual_seed(434) + s = torch.rand(n, 128, device="cuda", requires_grad=True) + bias = torch.randn(128, device="cuda") * 0.01 + ids, weights = route_scores_cuda(s, bias) + rid, rw = route_scores(s, bias) + assert torch.equal(ids, rid) + torch.testing.assert_close(weights, rw, atol=2e-7, rtol=2e-6) + g = torch.randn_like(weights) + actual = torch.autograd.grad(weights, s, g, retain_graph=True)[0] + expected = torch.autograd.grad(rw, s, g)[0] + torch.testing.assert_close(actual, expected, atol=2e-6, rtol=2e-5) + if n: + row = n // 2 + one = s[row : row + 1].detach().clone().requires_grad_() + oi, ow = route_scores_cuda(one, bias) + og = torch.autograd.grad(ow, one, g[row : row + 1])[0] + assert torch.equal(oi[0], ids[row]) + assert torch.equal(ow.view(torch.int32)[0], weights.view(torch.int32)[row]) + assert torch.equal(og.view(torch.int32)[0], actual.view(torch.int32)[row]) + + +def test_cuda_ties_and_zero(): + pytest.importorskip("triton") + from rl_engine.kernels.ops.triton.nemotron_router import route_scores_cuda + + for value in (0.0, 0.5): + scores = torch.full((17, 128), value, device="cuda", requires_grad=True) + ids, weights = route_scores_cuda(scores, torch.zeros(128, device="cuda")) + assert torch.equal(ids, torch.arange(6, device="cuda").expand(17, -1)) + assert torch.isfinite(torch.autograd.grad(weights.sum(), scores)[0]).all() + + +def test_cuda_near_ties_bias_and_invalid_values(): + from rl_engine.kernels.ops.triton.nemotron_router import route_scores_cuda + + scores = torch.full((1, 128), 0.5, device="cuda") + scores[0, 7:9] = torch.nextafter(scores[0, 7:9], torch.ones(2, device="cuda")) + bias = torch.zeros(128, device="cuda") + bias[80] = 1 + ids, weights = route_scores_cuda(scores, bias) + assert ids.tolist() == [[0, 1, 2, 7, 8, 80]] + # The large correction selects expert 80 but must not boost its weight. + assert weights[0, -1] == weights[0, 0] + for value in (float("nan"), float("inf"), -1.0, 2.0): + bad = scores.clone() + bad[0, 0] = value + with pytest.raises(ValueError, match="finite sigmoid"): + route_scores_cuda(bad, bias) + + +def _reference(x, w, b): + s = (x.double() @ w.double().t()).sigmoid() + ids = torch.argsort(s + b.double(), descending=True, stable=True)[:, :6] + ids = ids.sort(1).values + values = s.gather(1, ids) + weights = values / (values.sum(1, keepdim=True) + 1e-20) * 2.5 + perm = torch.argsort(ids.flatten(), stable=True) + offsets = torch.cat((ids.new_zeros(1), torch.bincount(ids.flatten(), minlength=128).cumsum(0))) + return ids, weights, perm, offsets, x.double()[perm // 6] + + +def _bits(a, b): + assert torch.equal(a.contiguous().view(torch.uint8), b.contiguous().view(torch.uint8)) + + +@_SM90 +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("n", [0, 1, 7, 17, 129, 513, 4096]) +def test_full_reference_and_batch_invariance(n, dtype): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + torch.manual_seed(434) + x = torch.randn(n, 2688, device="cuda", dtype=dtype).requires_grad_() + w = (torch.randn(128, 2688, device="cuda") * 0.01).requires_grad_() + b = torch.randn(128, device="cuda") * 0.001 + result = nemotron_router_cuda(x, w, b) + xr = x.detach().double().requires_grad_() + wr = w.detach().double().requires_grad_() + ref = _reference(xr, wr, b) + for i in (0, 2, 3): + assert torch.equal(result[i], ref[i]) + torch.testing.assert_close(result[1].double(), ref[1], atol=5e-7, rtol=3e-6) + assert torch.equal(result[4].double(), ref[4]) + gw = torch.randn_like(result[1]) + gp = torch.randn_like(result[4]) * 0.1 + dx, dw = torch.autograd.grad((result[1], result[4]), (x, w), (gw, gp)) + drx, drw = torch.autograd.grad((ref[1], ref[4]), (xr, wr), (gw.double(), gp.double())) + # Each BF16 input-gradient branch rounds before their autograd addition. + torch.testing.assert_close( + dx.double(), + drx, + atol=0.008 if dtype == torch.bfloat16 else 2e-6, + rtol=0.012 if dtype == torch.bfloat16 else 3e-5, + ) + torch.testing.assert_close(dw.double(), drw, atol=2e-5, rtol=3e-4) + if n: + inv = torch.argsort(result[2]) + for token in sorted({0, n // 2, n - 1}): + xx = x.detach()[token : token + 1].clone().requires_grad_() + one = nemotron_router_cuda(xx, w, b) + _bits(result[0][token : token + 1], one[0]) + _bits(result[1][token : token + 1], one[1]) + gp_token = gp[inv[token * 6 : (token + 1) * 6]] + da = torch.autograd.grad((one[1], one[4]), xx, (gw[token : token + 1], gp_token))[0] + _bits(dx[token : token + 1], da) + repeat = nemotron_router_cuda(x, w, b) + dx2, dw2 = torch.autograd.grad((repeat[1], repeat[4]), (x, w), (gw, gp)) + _bits(dx, dx2) + _bits(dw, dw2) + + +@_SM90 +def test_full_ties_and_payload_only_backward(): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + x = torch.randn(43, 2688, device="cuda", requires_grad=True) + w = torch.zeros(128, 2688, device="cuda", requires_grad=True) + b = torch.zeros(128, device="cuda") + ids, weights, perm, offsets, packed = nemotron_router_cuda(x, w, b) + assert torch.equal(ids, torch.arange(6, device="cuda").expand(43, -1)) + assert torch.equal(perm, torch.arange(258, device="cuda").reshape(43, 6).t().flatten()) + assert torch.equal(offsets[:7], torch.arange(7, device="cuda") * 43) + assert (offsets[7:] == 258).all() + dx = torch.autograd.grad(packed.sum(), x)[0] + assert torch.equal(dx, torch.full_like(x, 6)) + # Weight-only backward must work when the dispatch output is unused. + result = nemotron_router_cuda(x, w, b) + dx, dw = torch.autograd.grad(result[1].sum(), (x, w)) + assert torch.isfinite(dx).all() and torch.isfinite(dw).all() + + +@_SM90 +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("branch", ["route", "payload", "both"]) +def test_joint_backward_preserves_branch_rounding(dtype, branch): + from rl_engine.kernels.ops.triton.nemotron_router import ( + _Dispatch, + _ProjectRoute, + nemotron_router_cuda, + ) + + torch.manual_seed(5434) + x = torch.randn(43, 2688, device="cuda", dtype=dtype).requires_grad_(branch != "route") + w = (torch.randn(128, 2688, device="cuda") * 0.02).requires_grad_(branch != "payload") + bias = torch.zeros(128, device="cuda") + gw = torch.randn(43, 6, device="cuda") + gp = torch.randn(43 * 6, 2688, device="cuda", dtype=dtype) + actual = nemotron_router_cuda(x, w, bias) + ids, weights = _ProjectRoute.apply(x, w, bias) + perm, offsets, packed = _Dispatch.apply(x, ids) + separate = (ids, weights, perm, offsets, packed) + inputs = tuple(v for v in (x, w) if v.requires_grad) + indices = (1, 4) if branch == "both" else ((1,) if branch == "route" else (4,)) + cots = (gw, gp) if branch == "both" else ((gw,) if branch == "route" else (gp,)) + da = torch.autograd.grad(tuple(actual[i] for i in indices), inputs, cots) + ds = torch.autograd.grad(tuple(separate[i] for i in indices), inputs, cots) + for a, s in zip((*actual, *da), (*separate, *ds)): + _bits(a, s) + + +@_SM90 +def test_unsupported_geometry_fails(): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + with pytest.raises(ValueError): + nemotron_router_cuda( + torch.empty(3, 128, device="cuda"), + torch.empty(128, 128, device="cuda"), + torch.empty(128, device="cuda"), + ) + + +@_SM90 +def test_registry_real_provider_forward_backward(): + from rl_engine.kernels.ops.triton.nemotron_router import NemotronRouterOp + from rl_engine.kernels.registry import KernelRegistry + + registry = KernelRegistry() + op = registry.get_op("nemotron_router_dispatch", device="cuda") + assert isinstance(op, NemotronRouterOp) + x = torch.zeros(2, 2688, device="cuda", requires_grad=True) + w = torch.zeros(128, 2688, device="cuda", requires_grad=True) + b = torch.zeros(128, device="cuda") + ids, weights, perm, offsets, packed = op(x, w, b) + assert torch.equal(ids, torch.arange(6, device="cuda").expand(2, -1)) + assert offsets[-1].item() == 12 + dx, dw = torch.autograd.grad((weights.sum() + packed.sum()), (x, w)) + assert torch.equal(dx, torch.full_like(x, 6)) + assert torch.equal(dw, torch.zeros_like(w)) + + +@_SM90 +@pytest.mark.parametrize("n", [1, 42, 43, 85, 86, 8193, 65536]) +@pytest.mark.parametrize("concentrated", [False, True]) +def test_integer_dispatch_permutation_boundaries(n, concentrated): + """Independent stable-sort oracle at either side of 256-route block tails.""" + import triton + + from rl_engine.kernels.ops.triton.nemotron_router import _dispatch_counts, _dispatch_map + + torch.manual_seed(434 + n) + ids = torch.rand(n, 128, device="cuda").argsort(1)[:, :6].sort(1).values.contiguous() + if concentrated: + # Include the largest legal expert, adjacent to the masked sentinel. + ids[:] = torch.tensor([0, 1, 2, 125, 126, 127], device="cuda") + routes = n * 6 + blocks = triton.cdiv(routes, 256) + counts = torch.empty((128, blocks), device="cuda", dtype=torch.int64) + _dispatch_counts[(blocks,)](ids, counts, n, blocks) + prefix = counts.cumsum(1) + offsets = torch.cat((ids.new_zeros(1), prefix[:, -1].cumsum(0))) + permutation = torch.full((routes,), -1, device="cuda", dtype=torch.int64) + inverse = torch.full_like(permutation, -1) + _dispatch_map[(blocks,)](ids, prefix, offsets, permutation, inverse, n, blocks) + expected = torch.argsort(ids.flatten(), stable=True) + assert torch.equal(permutation, expected) + assert torch.equal(inverse[permutation], torch.arange(routes, device="cuda")) + expected_counts = torch.bincount(ids.flatten(), minlength=128) + assert torch.equal(offsets[1:] - offsets[:-1], expected_counts) + + +@_SM90 +@pytest.mark.parametrize("seed", [0, 17, 2026]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("n", [1, 33, 257, 8193]) +def test_row_relocation_and_gradient(seed, dtype, n): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + torch.manual_seed(seed) + x = torch.randn(n, 2688, device="cuda", dtype=dtype).requires_grad_() + w = torch.randn(128, 2688, device="cuda") * 0.01 + bias = torch.randn(128, device="cuda") * 0.001 + out = nemotron_router_cuda(x, w, bias) + gradient = torch.randn_like(out[1]) + dx = torch.autograd.grad(out[1], x, gradient)[0] + # Reverse row order moves rows between fixed output tiles and tile lanes. + xx = x.detach().flip(0).contiguous().requires_grad_() + other = nemotron_router_cuda(xx, w, bias) + _bits(out[0].flip(0), other[0]) + _bits(out[1].flip(0), other[1]) + dx_other = torch.autograd.grad(other[1], xx, gradient.flip(0).contiguous())[0] + _bits(dx.flip(0), dx_other) + for index in sorted({0, min(n - 1, 31), min(n - 1, 32), n - 1}): + row = x.detach()[index : index + 1].clone().requires_grad_() + single = nemotron_router_cuda(row, w, bias) + _bits(out[0][index : index + 1], single[0]) + _bits(out[1][index : index + 1], single[1]) + drow = torch.autograd.grad(single[1], row, gradient[index : index + 1])[0] + _bits(dx[index : index + 1], drow) + # Every flat route occurs exactly once, including multi-block expert segments. + _bits(out[2].sort().values, torch.arange(n * 6, device="cuda")) + assert out[3][0].item() == 0 and out[3][-1].item() == n * 6 + assert (out[3][1:] >= out[3][:-1]).all() + _bits(out[4], x.detach()[out[2] // 6]) + + +@_SM90 +@pytest.mark.parametrize("scale", [1e-8, 1.0, 1e8]) +def test_projection_cancellation_and_dynamic_range(scale): + from rl_engine.kernels.ops.triton.nemotron_router import _Projection + + # Products cancel in adjacent pairs. Perturbations prevent a trivial zero-only test. + torch.manual_seed(912) + x = torch.ones(65, 2688, device="cuda") * scale + x[:, ::2] *= -1 + w = torch.randn(128, 1344, device="cuda").repeat_interleave(2, dim=1) + w[:, 0] += 0.125 + result = _Projection.apply(x, w) + reference = x.double() @ w.double().t() + torch.testing.assert_close(result.double(), reference, atol=3e-5 * scale, rtol=3e-4) + _bits(result[32:33], _Projection.apply(x[32:33].contiguous(), w)) + + +@_SM90 +@pytest.mark.parametrize("n", [1, 129]) +def test_graph_replay_observes_input_and_weight_updates(n): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + torch.manual_seed(45) + x = torch.randn(n, 2688, device="cuda", dtype=torch.bfloat16) + w = torch.randn(128, 2688, device="cuda") * 0.01 + bias = torch.zeros(128, device="cuda") + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + for _ in range(3): + nemotron_router_cuda(x, w, bias) + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = nemotron_router_cuda(x, w, bias) + for _ in range(3): + x.normal_() + w.normal_(std=0.02) + bias.normal_(std=0.001) + graph.replay() + expected = nemotron_router_cuda(x, w, bias) + for actual, reference in zip(captured, expected): + _bits(actual, reference) + + +@_SM90 +def test_outer_autocast_preserves_fp32_contract(): + from rl_engine.kernels.ops.triton.nemotron_router import nemotron_router_cuda + + torch.manual_seed(9) + x = torch.randn(33, 2688, device="cuda", requires_grad=True) + w = (torch.randn(128, 2688, device="cuda") * 0.01).requires_grad_() + bias = torch.zeros(128, device="cuda") + baseline = nemotron_router_cuda(x, w, bias) + with torch.autocast("cuda", dtype=torch.bfloat16): + actual = nemotron_router_cuda(x, w, bias) + for a, b in zip(actual, baseline): + _bits(a, b) + g = torch.randn_like(actual[1]) + ga = torch.autograd.grad(actual[1], (x, w), g) + gb = torch.autograd.grad(baseline[1], (x, w), g) + for a, b in zip(ga, gb): + _bits(a, b) + + +@_SM90 +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) +def test_maximum_batch_last_payload_and_gradient(dtype): + from rl_engine.kernels.ops.triton.nemotron_router import MAX_TOKENS, nemotron_router_cuda + + torch.manual_seed(734) + x = torch.randn(MAX_TOKENS, 2688, device="cuda", dtype=dtype).requires_grad_() + w = torch.randn(128, 2688, device="cuda") * 0.01 + bias = torch.zeros(128, device="cuda") + output = nemotron_router_cuda(x, w, bias) + # Check every packed row, including addresses near the supported index bound. + _bits(output[4], x.detach()[output[2] // 6]) + row = x.detach()[-1:].clone().requires_grad_() + single = nemotron_router_cuda(row, w, bias) + _bits(output[0][-1:], single[0]) + _bits(output[1][-1:], single[1]) + upstream = torch.zeros_like(output[1]) + upstream[-1] = torch.arange(6, device="cuda") + dx = torch.autograd.grad(output[1], x, upstream)[0] + single_dx = torch.autograd.grad(single[1], row, upstream[-1:])[0] + _bits(dx[-1:], single_dx) + assert not torch.count_nonzero(dx[:-1]) + + +@_SM90 +def test_excess_tokens_rejected_before_allocation(): + from rl_engine.kernels.ops.triton.nemotron_router import MAX_TOKENS, nemotron_router_cuda + + # A broadcast view exercises the metadata guard without allocating a huge input. + x = torch.zeros(1, device="cuda").expand(MAX_TOKENS + 1, 2688) + w = torch.zeros(128, 2688, device="cuda") + bias = torch.zeros(128, device="cuda") + with pytest.raises(ValueError, match="at most 65536 tokens"): + nemotron_router_cuda(x, w, bias) + + +@_SM90 +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("concentrated", [False, True]) +def test_registry_nonlinear_consumer_forward_backward(dtype, concentrated): + from rl_engine.kernels.registry import KernelRegistry + + torch.manual_seed(4434) + x = torch.randn(33, 2688, device="cuda", dtype=dtype).requires_grad_() + w = (torch.randn(128, 2688, device="cuda") * 0.01).requires_grad_() + bias = torch.randn(128, device="cuda") * 0.01 + if concentrated: + bias.zero_() + bias[:6] = 4 + scale = torch.linspace(0.5, 1.5, 128, device="cuda") + op = KernelRegistry().get_op("nemotron_router_dispatch", device=x.device) + ids, weights, permutation, offsets, packed = op(x, w, bias) + assert ids.dtype == permutation.dtype == offsets.dtype == torch.int64 + assert weights.dtype == torch.float32 and packed.dtype == dtype + spans = offsets.tolist() # Correctness-only consumer; not a benchmark. + chunks = [ + torch.relu(packed[spans[e] : spans[e + 1]].float() * scale[e] + 0.01).square() + for e in range(128) + ] + expert_output = torch.cat(chunks) + slots = expert_output[permutation.argsort()].reshape(33, 6, 2688) + actual = (slots * weights[:, :, None]).sum(1) + x.float().tanh() + + # Independent token-major FP64 expression never uses packed data or offsets. + xr, wr = x.detach().double().requires_grad_(), w.detach().double().requires_grad_() + scores = (xr @ wr.t()).sigmoid() + gold_ids = torch.argsort(scores + bias.double(), descending=True, stable=True)[:, :6] + gold_ids = gold_ids.sort(1).values + assert torch.equal(ids, gold_ids) + values = scores.gather(1, gold_ids) + gold_weights = values / (values.sum(1, keepdim=True) + 1e-20) * 2.5 + expert_values = torch.relu(xr[:, None, :] * scale.double()[gold_ids, None] + 0.01).square() + expected = (expert_values * gold_weights[:, :, None]).sum(1) + xr.tanh() + g = torch.randn_like(actual) + dx, dw = torch.autograd.grad(actual, (x, w), g) + rx, rw = torch.autograd.grad(expected, (xr, wr), g.double()) + torch.testing.assert_close(actual.double(), expected, atol=2e-5, rtol=3e-6) + torch.testing.assert_close(dw.double(), rw, atol=5e-5, rtol=3e-4) + error = (dx.double() - rx).norm() / rx.norm() + assert error < (0.01 if dtype == torch.bfloat16 else 1e-5) diff --git a/tests/nemotron/test_nemotron_router_reference.py b/tests/nemotron/test_nemotron_router_reference.py new file mode 100644 index 000000000..e14ed53ad --- /dev/null +++ b/tests/nemotron/test_nemotron_router_reference.py @@ -0,0 +1,129 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors +"""Independent scalar checks and differentiability of the draft reference.""" + +import math + +import pytest +import torch + +from rl_engine.kernels.ops.pytorch.nemotron_router import ( + dispatch_tokens, + route_scores, + router_reference, +) + + +def scalar_oracle(x, w, bias, k): + ids, weights = [], [] + for row in x.tolist(): + logits = [math.fsum(a * b for a, b in zip(row, wr)) for wr in w.tolist()] + scores = [1 / (1 + math.exp(-v)) for v in logits] + chosen = sorted(sorted(range(len(scores)), key=lambda e: (-(scores[e] + bias[e]), e))[:k]) + denom = math.fsum(scores[e] for e in chosen) + 1e-20 + ids.append(chosen) + weights.append([2.5 * scores[e] / denom for e in chosen]) + return torch.tensor(ids), torch.tensor(weights, dtype=torch.float64) + + +def test_independent_scalar_and_dispatch(): + g = torch.Generator().manual_seed(434) + x = torch.randn(7, 9, generator=g, dtype=torch.float64) + w = torch.randn(8, 9, generator=g, dtype=torch.float64) + b = torch.linspace(-0.2, 0.2, 8, dtype=torch.float64) + result = router_reference(x, w, b, top_k=3) + ids, weights = scalar_oracle(x, w, b.tolist(), 3) + assert torch.equal(result.expert_ids, ids) + torch.testing.assert_close(result.weights, weights, atol=1e-14, rtol=1e-14) + slots = [(e, t, s) for t, row in enumerate(ids.tolist()) for s, e in enumerate(row)] + slots.sort() + assert result.permutation.tolist() == [t * 3 + s for e, t, s in slots] + assert torch.equal(result.packed_tokens, x[[t for e, t, s in slots]]) + for e in range(8): + lo, hi = result.expert_offsets[e : e + 2].tolist() + assert hi - lo == sum(v[0] == e for v in slots) + + +def test_ties_choose_lowest_ids(): + ids, weights = route_scores(torch.full((3, 128), 0.5), torch.zeros(128)) + assert torch.equal(ids, torch.arange(6).expand(3, -1)) + torch.testing.assert_close(weights.sum(1), torch.full((3,), 2.5)) + + +def test_bias_is_selection_only(): + ids, weights = route_scores(torch.tensor([[0.1, 0.2, 0.7]]), torch.tensor([1.0, 1.0, 0.0]), 2) + assert ids.tolist() == [[0, 1]] + torch.testing.assert_close(weights, torch.tensor([[2.5 / 3, 5.0 / 3]])) + + +def test_zero_scores_and_empty_tokens(): + ids, weights = route_scores(torch.zeros(1, 128), torch.zeros(128)) + assert torch.equal(weights, torch.zeros_like(weights)) + result = router_reference(torch.empty(0, 4), torch.zeros(128, 4), torch.zeros(128)) + assert result.packed_tokens.shape == (0, 4) + assert result.expert_offsets.tolist() == [0] * 129 + + +def test_score_batch_position_invariance(): + scores = torch.rand(17, 128, generator=torch.Generator().manual_seed(1)) + bias = torch.zeros(128) + ids, weights = route_scores(scores, bias) + for t in [0, 7, 16]: + one_ids, one_weights = route_scores(scores[t : t + 1], bias) + assert torch.equal(one_ids[0], ids[t]) + assert torch.equal(one_weights[0], weights[t]) + + +def test_gradcheck_weights_and_payload(): + g = torch.Generator().manual_seed(4) + x = (torch.randn(3, 4, generator=g, dtype=torch.float64) * 0.1).requires_grad_() + w = (torch.randn(8, 4, generator=g, dtype=torch.float64) * 0.1).requires_grad_() + # Well-separated selection keeps finite differences away from top-k discontinuities. + bias = torch.arange(8, dtype=torch.float64) * 0.2 + + def fn(a, b): + r = router_reference(a, b, bias, top_k=3) + return r.weights, r.packed_tokens + + assert torch.autograd.gradcheck(fn, (x, w)) + + +def test_payload_gradient_counts_all_routes(): + x = torch.randn(2, 3, requires_grad=True) + _, _, packed = dispatch_tokens(x, torch.tensor([[0, 2], [1, 2]]), 4) + packed.sum().backward() + assert torch.equal(x.grad, torch.full_like(x, 2)) + + +@pytest.mark.parametrize("bad", [float("nan"), float("inf"), -0.1, 1.1]) +def test_reject_invalid_scores(bad): + with pytest.raises(ValueError): + route_scores(torch.tensor([[bad, 0.5]]), torch.zeros(2), 1) + + +def test_reject_duplicate_experts(): + with pytest.raises(ValueError): + dispatch_tokens(torch.zeros(1, 3), torch.tensor([[1, 1]]), 2) + + +def test_registry_fails_closed_and_selects_only_router(monkeypatch): + from rl_engine.kernels.registry import KernelRegistry, OpBackend + + registry = KernelRegistry() + for device in ("cpu", "mps"): + with pytest.raises(RuntimeError, match="SM90"): + registry.get_op("nemotron_router_dispatch", device=device) + sentinel = object() + calls = [] + + def load(backend): + calls.append(backend) + return sentinel + + monkeypatch.setattr(registry, "_get_or_create_backend", load) + monkeypatch.setattr(registry, "_platform_for_device", lambda device: "cuda") + assert registry.get_op("nemotron_router_dispatch", "cuda") is sentinel + assert calls == [OpBackend.TRITON_NEMOTRON_ROUTER] + monkeypatch.setattr(registry, "_get_or_create_backend", lambda backend: None) + with pytest.raises(RuntimeError, match="No functional backend"): + registry.get_op("nemotron_router_dispatch", "cuda")