From 29e360b3c9e932a06b2431d6978dca41fbfe8c32 Mon Sep 17 00:00:00 2001 From: a-kaa Date: Mon, 22 Jun 2026 00:35:38 +0800 Subject: [PATCH 1/5] Add PyTorch matmul reference operator --- docs/operators/README.md | 1 + docs/operators/matmul.md | 96 ++++++++++++ .../kernels/ops/pytorch/linear/__init__.py | 6 + .../kernels/ops/pytorch/linear/matmul.py | 31 ++++ rl_engine/kernels/registry.py | 4 + tests/test_matmul.py | 146 ++++++++++++++++++ 6 files changed, 284 insertions(+) create mode 100644 docs/operators/matmul.md create mode 100644 rl_engine/kernels/ops/pytorch/linear/__init__.py create mode 100644 rl_engine/kernels/ops/pytorch/linear/matmul.py create mode 100644 tests/test_matmul.py diff --git a/docs/operators/README.md b/docs/operators/README.md index c4eae603f..5a6d51848 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -22,5 +22,6 @@ Every operator page should include: - [Fused Linear LogP](linear-logp.md) - [GRPO Loss](grpo-loss.md) - [Policy Ratio + KL Penalty](ratio-kl.md) +- [Matmul](matmul.md) - [Sampling](sampling.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/matmul.md b/docs/operators/matmul.md new file mode 100644 index 000000000..eaef492c8 --- /dev/null +++ b/docs/operators/matmul.md @@ -0,0 +1,96 @@ +# Matmul + +Matmul provides the PyTorch reference GEMM operator. It targets dense projection shapes used in Qwen3-style post-training workloads, including Q, K, V, O, gate, up, and down projections. + +This page documents the PyTorch baseline version. + +## Entry Point + +```python +from rl_engine.kernels.registry import kernel_registry + +matmul = kernel_registry.get_op("matmul") +output = matmul.forward(a, b) +reference = matmul.forward_fp32(a, b) +``` + +The operator can also be imported directly: + +```python +from rl_engine.kernels.ops.pytorch.linear import NativeMatmulOp + +matmul = NativeMatmulOp() +``` + +## Backend + +| Backend | Wrapper | Native symbol | Notes | +| --- | --- | --- | --- | +| PyTorch native | `NativeMatmulOp` | None | Reference baseline; no split-K or manual blocking. | + +`kernel_registry.get_op("matmul")` dispatches to the PyTorch native backend on CPU, +CUDA, and ROCm. CUDA/Triton fused GEMM kernels should compare against this reference. + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `a` | `[..., K]` | `float32`, `bfloat16`, or `float16` | Left operand. | +| `b` | `[K, N]` | Same floating dtype family | Right operand in `[in, out]` layout. | +| Output | `[..., N]` | See below | Result of `a @ b`. | + +`forward(...)` returns `a.dtype`. `forward_fp32(...)` casts both inputs to +`float32`, calls `torch.matmul` once, and returns `float32`. + +## Qwen3 Projection Shapes + +The reference covers the Qwen3-8B projection dimensions: + +| Projection | Shape | +| --- | --- | +| `q_proj`, `o_proj` | `4096 -> 4096` | +| `k_proj`, `v_proj` | `4096 -> 1024` | +| `gate_proj`, `up_proj` | `4096 -> 12288` | +| `down_proj` | `12288 -> 4096` | + +LM head uses a separate operator because its weight follows the HF `[out, in]` +layout and is internally transposed. + +## Reference Semantics + +```python +ref = torch.matmul(a.float(), b.float()) +``` + +The gold path intentionally avoids split-K, tiled reductions, or any manual +accumulation order. Those optimizations belong in downstream fused kernels and +should be validated against this baseline. + +## Accuracy + +Matmul is categorized as a `reduction` operator in the numerical contract. +Expected comparison behavior: + +| Path | Expected dtype | Purpose | +| --- | --- | --- | +| `forward` | `a.dtype` | Candidate dtype behavior. | +| `forward_fp32` | `torch.float32` | Deterministic reference output. | + +Batch invariance is expected for the reference path: applying matmul to a full +batch and then slicing a row should match applying matmul to that row alone. + +## Tests + +```bash +python -m pytest tests/test_matmul.py -q +``` + +The test covers shape, dtype behavior, fp32 reference equivalence, batch +invariance, Qwen3 projection dimensions, and registry dispatch. + +## Implementation Files + +- `rl_engine/kernels/ops/pytorch/linear/matmul.py` +- `rl_engine/kernels/ops/pytorch/linear/__init__.py` +- `rl_engine/kernels/registry.py` +- `tests/test_matmul.py` diff --git a/rl_engine/kernels/ops/pytorch/linear/__init__.py b/rl_engine/kernels/ops/pytorch/linear/__init__.py new file mode 100644 index 000000000..0219e739b --- /dev/null +++ b/rl_engine/kernels/ops/pytorch/linear/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +from rl_engine.kernels.ops.pytorch.linear.matmul import NativeMatmulOp + +__all__ = ["NativeMatmulOp"] diff --git a/rl_engine/kernels/ops/pytorch/linear/matmul.py b/rl_engine/kernels/ops/pytorch/linear/matmul.py new file mode 100644 index 000000000..8aa5c122f --- /dev/null +++ b/rl_engine/kernels/ops/pytorch/linear/matmul.py @@ -0,0 +1,31 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +from __future__ import annotations + +import torch +from torch import Tensor + + +class NativeMatmulOp: + """Pure PyTorch reference GEMM. + + It intentionally uses one `torch.matmul` call in fp32 for + the gold path and does not implement split-K or manual blocked accumulation. + """ + + op_class = "reduction" + + def __init__(self) -> None: + pass + + def __call__(self, a: Tensor, b: Tensor) -> Tensor: + return self.forward(a, b) + + def forward(self, a: Tensor, b: Tensor) -> Tensor: + """Compute `a @ b` and return the input dtype.""" + return self.forward_fp32(a, b).to(dtype=a.dtype) + + def forward_fp32(self, a: Tensor, b: Tensor) -> Tensor: + """fp32 gold standard: cast inputs to fp32, then call `torch.matmul` once.""" + return torch.matmul(a.float(), b.float()) diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 678015713..3bd592c6a 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -53,6 +53,7 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): TRITON_GENERIC = "rl_engine.kernels.ops.triton.generic.TritonOp" PYTORCH_ATTN = "rl_engine.kernels.ops.pytorch.attention.NativeAttentionOp" PYTORCH_NATIVE = "rl_engine.kernels.ops.pytorch.loss.logp.NativeLogpOp" + PYTORCH_NATIVE_MATMUL = "rl_engine.kernels.ops.pytorch.linear.matmul.NativeMatmulOp" class KernelRegistry: @@ -90,6 +91,7 @@ def __init__(self): "linear_logp": [OpBackend.TRITON_LINEAR_LOGP, OpBackend.PYTORCH_LINEAR_LOGP], "ratio_kl": [OpBackend.TRITON_RATIO_KL, OpBackend.PYTORCH_RATIO_KL], # Default dispatch logic for new operators + "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], }, "rocm": { "logp": [OpBackend.ROCM_AITER, OpBackend.TRITON_GENERIC, OpBackend.PYTORCH_NATIVE], @@ -101,6 +103,7 @@ def __init__(self): "grpo_loss": [OpBackend.TRITON_GRPO_LOSS, OpBackend.PYTORCH_GRPO_LOSS], "linear_logp": [OpBackend.TRITON_LINEAR_LOGP, OpBackend.PYTORCH_LINEAR_LOGP], "ratio_kl": [OpBackend.TRITON_RATIO_KL, OpBackend.PYTORCH_RATIO_KL], + "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], }, "cpu": { "logp": [OpBackend.PYTORCH_NATIVE], @@ -108,6 +111,7 @@ def __init__(self): "grpo_loss": [OpBackend.PYTORCH_GRPO_LOSS], "linear_logp": [OpBackend.PYTORCH_LINEAR_LOGP], "ratio_kl": [OpBackend.PYTORCH_RATIO_KL], + "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], }, } logger.info(f"KernelRegistry initialized for {device_ctx.device_type}") diff --git a/tests/test_matmul.py b/tests/test_matmul.py new file mode 100644 index 000000000..fb2bfda17 --- /dev/null +++ b/tests/test_matmul.py @@ -0,0 +1,146 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Tests for NativeMatmulOp, the PyTorch fp32 GEMM reference.""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.kernels.ops.pytorch.linear.matmul import NativeMatmulOp + + +QWEN3_HIDDEN = 4096 +QWEN3_INTERMEDIATE = 12288 + + +def _make_inputs( + batch: int, + seq: int, + k: int, + n: int, + *, + dtype: torch.dtype = torch.float32, + seed: int = 123, +) -> tuple[torch.Tensor, torch.Tensor]: + gen = torch.Generator().manual_seed(seed) + a = torch.randn(batch, seq, k, generator=gen, dtype=dtype) + b = torch.randn(k, n, generator=gen, dtype=dtype) + return a, b + + +class TestNativeMatmulOpCorrectness: + def test_output_shape_matches_matmul_contract(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32) + out = op.forward_fp32(a, b) + assert out.shape == (2, 16, 32) + + def test_forward_fp32_returns_fp32(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32, dtype=torch.bfloat16) + out = op.forward_fp32(a, b) + assert out.dtype == torch.float32 + + @pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) + def test_forward_returns_input_dtype(self, dtype): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32, dtype=dtype) + out = op.forward(a, b) + assert out.dtype == dtype + + def test_call_equals_forward(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32) + assert torch.equal(op(a, b), op.forward(a, b)) + + def test_matches_single_fp32_torch_matmul(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32) + out = op.forward_fp32(a, b) + ref = torch.matmul(a.float(), b.float()) + assert torch.equal(out, ref) + + def test_pure_function_no_inplace(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32) + a_orig = a.clone() + b_orig = b.clone() + _ = op.forward_fp32(a, b) + assert torch.equal(a, a_orig) + assert torch.equal(b, b_orig) + + def test_op_class_is_reduction(self): + assert NativeMatmulOp.op_class == "reduction" + + +class TestNativeMatmulOpBatchInvariance: + def test_batch1_vs_batchN_bitwise(self): + op = NativeMatmulOp() + a, b = _make_inputs(4, 16, 64, 32, seed=321) + full_out = op.forward_fp32(a, b) + for row in range(a.shape[0]): + single_out = op.forward_fp32(a[row : row + 1], b) + assert torch.equal(full_out[row], single_out[0]), ( + f"Batch invariance broken at row {row}" + ) + + def test_batch_invariance_with_padding(self): + op = NativeMatmulOp() + a_valid, b = _make_inputs(2, 16, 64, 32, seed=456) + gen = torch.Generator().manual_seed(789) + padding = torch.randn(3, 16, 64, generator=gen) + a_padded = torch.cat([a_valid, padding], dim=0) + out_valid = op.forward_fp32(a_valid, b) + out_padded = op.forward_fp32(a_padded, b) + assert torch.equal(out_valid[0], out_padded[0]) + assert torch.equal(out_valid[1], out_padded[1]) + + +class TestNativeMatmulOpAccuracy: + @pytest.mark.parametrize( + "dtype, atol, rtol", + [ + (torch.float32, 1e-4, 1e-4), + (torch.bfloat16, 5e-2, 2e-2), + (torch.float16, 1e-3, 1e-3), + ], + ) + def test_forward_vs_fp32_within_tolerance(self, dtype, atol, rtol): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32, dtype=dtype) + out_typed = op.forward(a, b).float() + out_fp32 = op.forward_fp32(a, b) + diff = (out_typed - out_fp32).abs().max().item() + assert torch.allclose(out_typed, out_fp32, atol=atol, rtol=rtol), ( + f"dtype={dtype}, max_abs_error={diff:.3e} exceeds " + f"atol={atol}, rtol={rtol}" + ) + + +class TestNativeMatmulOpQwen3Shapes: + @pytest.mark.parametrize( + "k, n, label", + [ + (QWEN3_HIDDEN, QWEN3_HIDDEN, "q_proj/o_proj"), + (QWEN3_HIDDEN, 1024, "k_proj/v_proj"), + (QWEN3_HIDDEN, QWEN3_INTERMEDIATE, "gate_proj/up_proj"), + (QWEN3_INTERMEDIATE, QWEN3_HIDDEN, "down_proj"), + ], + ) + def test_qwen3_projection_reduction_dims(self, k, n, label): + del label + op = NativeMatmulOp() + a, b = _make_inputs(1, 2, k, n, seed=42) + out = op.forward_fp32(a, b) + assert out.shape == (1, 2, n) + assert out.dtype == torch.float32 + + +class TestNativeMatmulOpRegistry: + def test_registry_returns_matmul_op(self): + from rl_engine.kernels.registry import kernel_registry + + op = kernel_registry.get_op("matmul") + assert isinstance(op, NativeMatmulOp) From f364d8fde6dfea433c1130fba53600edfee37dc7 Mon Sep 17 00:00:00 2001 From: a-kaa Date: Tue, 23 Jun 2026 00:04:10 +0800 Subject: [PATCH 2/5] format with pre-commit --- tests/test_matmul.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/tests/test_matmul.py b/tests/test_matmul.py index fb2bfda17..83d3627aa 100644 --- a/tests/test_matmul.py +++ b/tests/test_matmul.py @@ -10,7 +10,6 @@ from rl_engine.kernels.ops.pytorch.linear.matmul import NativeMatmulOp - QWEN3_HIDDEN = 4096 QWEN3_INTERMEDIATE = 12288 @@ -82,9 +81,9 @@ def test_batch1_vs_batchN_bitwise(self): full_out = op.forward_fp32(a, b) for row in range(a.shape[0]): single_out = op.forward_fp32(a[row : row + 1], b) - assert torch.equal(full_out[row], single_out[0]), ( - f"Batch invariance broken at row {row}" - ) + assert torch.equal( + full_out[row], single_out[0] + ), f"Batch invariance broken at row {row}" def test_batch_invariance_with_padding(self): op = NativeMatmulOp() @@ -114,8 +113,7 @@ def test_forward_vs_fp32_within_tolerance(self, dtype, atol, rtol): out_fp32 = op.forward_fp32(a, b) diff = (out_typed - out_fp32).abs().max().item() assert torch.allclose(out_typed, out_fp32, atol=atol, rtol=rtol), ( - f"dtype={dtype}, max_abs_error={diff:.3e} exceeds " - f"atol={atol}, rtol={rtol}" + f"dtype={dtype}, max_abs_error={diff:.3e} exceeds " f"atol={atol}, rtol={rtol}" ) From 64bd581a6c536c9712a96c5f121d0878428b9a94 Mon Sep 17 00:00:00 2001 From: a-kaa Date: Sat, 4 Jul 2026 14:53:09 +0800 Subject: [PATCH 3/5] Make native matmul a torch module --- rl_engine/kernels/ops/pytorch/linear/matmul.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/rl_engine/kernels/ops/pytorch/linear/matmul.py b/rl_engine/kernels/ops/pytorch/linear/matmul.py index 8aa5c122f..6455b2938 100644 --- a/rl_engine/kernels/ops/pytorch/linear/matmul.py +++ b/rl_engine/kernels/ops/pytorch/linear/matmul.py @@ -7,7 +7,7 @@ from torch import Tensor -class NativeMatmulOp: +class NativeMatmulOp(torch.nn.Module): """Pure PyTorch reference GEMM. It intentionally uses one `torch.matmul` call in fp32 for @@ -17,10 +17,7 @@ class NativeMatmulOp: op_class = "reduction" def __init__(self) -> None: - pass - - def __call__(self, a: Tensor, b: Tensor) -> Tensor: - return self.forward(a, b) + super().__init__() def forward(self, a: Tensor, b: Tensor) -> Tensor: """Compute `a @ b` and return the input dtype.""" From 65e2c238e4eb91481591bda6218d362aa80c5024 Mon Sep 17 00:00:00 2001 From: a-kaa Date: Sat, 4 Jul 2026 14:59:46 +0800 Subject: [PATCH 4/5] Add native matmul backward test --- tests/test_matmul.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/test_matmul.py b/tests/test_matmul.py index 83d3627aa..db237479f 100644 --- a/tests/test_matmul.py +++ b/tests/test_matmul.py @@ -74,6 +74,22 @@ def test_op_class_is_reduction(self): assert NativeMatmulOp.op_class == "reduction" +class TestNativeMatmulOpBackward: + def test_backward_matches_torch_matmul(self): + op = NativeMatmulOp() + a, b = _make_inputs(2, 16, 64, 32) + ref_a = a.clone().requires_grad_(True) + ref_b = b.clone().requires_grad_(True) + test_a = a.clone().requires_grad_(True) + test_b = b.clone().requires_grad_(True) + + torch.matmul(ref_a, ref_b).sum().backward() + op.forward_fp32(test_a, test_b).sum().backward() + + assert torch.equal(test_a.grad, ref_a.grad) + assert torch.equal(test_b.grad, ref_b.grad) + + class TestNativeMatmulOpBatchInvariance: def test_batch1_vs_batchN_bitwise(self): op = NativeMatmulOp() From 2a44ed72c4897b48f0d97941c98fe87d3ad75029 Mon Sep 17 00:00:00 2001 From: a-kaa Date: Sat, 4 Jul 2026 16:24:11 +0800 Subject: [PATCH 5/5] Add native matmul batch grad invariance test --- tests/test_matmul.py | 42 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 42 insertions(+) diff --git a/tests/test_matmul.py b/tests/test_matmul.py index db237479f..7885aa34d 100644 --- a/tests/test_matmul.py +++ b/tests/test_matmul.py @@ -5,6 +5,8 @@ from __future__ import annotations +from contextlib import contextmanager + import pytest import torch @@ -14,6 +16,16 @@ QWEN3_INTERMEDIATE = 12288 +@contextmanager +def _single_threaded_torch(): + old_threads = torch.get_num_threads() + torch.set_num_threads(1) + try: + yield + finally: + torch.set_num_threads(old_threads) + + def _make_inputs( batch: int, seq: int, @@ -112,6 +124,36 @@ def test_batch_invariance_with_padding(self): assert torch.equal(out_valid[0], out_padded[0]) assert torch.equal(out_valid[1], out_padded[1]) + def test_batch_grad_invariance(self): + op = NativeMatmulOp() + a, b = _make_inputs(4, 8, 512, 384, seed=654) + # Use a non-unit upstream gradient to exercise the real backward path. + grad_out = torch.randn(4, 8, 384, generator=torch.Generator().manual_seed(987)) + + with _single_threaded_torch(): + full_a = a.clone().requires_grad_(True) + full_b = b.clone().requires_grad_(True) + (op.forward_fp32(full_a, full_b) * grad_out).sum().backward() + + single_a_grads = [] + single_b_grads = [] + for row in range(a.shape[0]): + single_a = a[row : row + 1].clone().requires_grad_(True) + # The shared weight gradient is the sum of all per-batch contributions. + single_b = b.clone().requires_grad_(True) + single_grad_out = grad_out[row : row + 1] + (op.forward_fp32(single_a, single_b) * single_grad_out).sum().backward() + single_a_grads.append(single_a.grad[0]) + single_b_grads.append(single_b.grad) + + assert torch.equal(full_a.grad, torch.stack(single_a_grads)) + torch.testing.assert_close( + full_b.grad, + torch.stack(single_b_grads).sum(dim=0), + atol=1e-5, + rtol=1e-6, + ) + class TestNativeMatmulOpAccuracy: @pytest.mark.parametrize(