diff --git a/docs/operators/README.md b/docs/operators/README.md index 07e6dc853..c6f321540 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -26,6 +26,7 @@ Every operator page should include: - [RoPE](rope.md) - [LM Head](lm_head.md) - [Policy Ratio + KL Penalty](ratio-kl.md) +- [Matmul](matmul.md) - [Sampling](sampling.md) - [Token Embedding](embedding.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 index 86cf4c9da..0219e739b 100644 --- a/rl_engine/kernels/ops/pytorch/linear/__init__.py +++ b/rl_engine/kernels/ops/pytorch/linear/__init__.py @@ -1,2 +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..6455b2938 --- /dev/null +++ b/rl_engine/kernels/ops/pytorch/linear/matmul.py @@ -0,0 +1,28 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +from __future__ import annotations + +import torch +from torch import Tensor + + +class NativeMatmulOp(torch.nn.Module): + """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: + super().__init__() + + 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 c64720f9c..d6dda2f56 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -56,6 +56,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" PYTORCH_NATIVE_ROPE = "rl_engine.kernels.ops.pytorch.rotary_embedding.rope.NativeRoPEOp" PYTORCH_NATIVE_SILU = "rl_engine.kernels.ops.pytorch.activation.swiglu.NativeSiLUOp" PYTORCH_NATIVE_SWIGLU = "rl_engine.kernels.ops.pytorch.activation.swiglu.NativeSwiGLUOp" @@ -106,6 +107,7 @@ def __init__(self): "silu": [OpBackend.PYTORCH_NATIVE_SILU], "swiglu": [OpBackend.PYTORCH_NATIVE_SWIGLU], # Default dispatch logic for new operators + "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rope": [OpBackend.PYTORCH_NATIVE_ROPE], }, "rocm": { @@ -119,6 +121,7 @@ def __init__(self): "rope": [OpBackend.PYTORCH_NATIVE_ROPE], "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], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], @@ -132,6 +135,7 @@ def __init__(self): "rope": [OpBackend.PYTORCH_NATIVE_ROPE], "linear_logp": [OpBackend.PYTORCH_LINEAR_LOGP], "ratio_kl": [OpBackend.PYTORCH_RATIO_KL], + "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], diff --git a/tests/test_matmul.py b/tests/test_matmul.py new file mode 100644 index 000000000..7885aa34d --- /dev/null +++ b/tests/test_matmul.py @@ -0,0 +1,202 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Tests for NativeMatmulOp, the PyTorch fp32 GEMM reference.""" + +from __future__ import annotations + +from contextlib import contextmanager + +import pytest +import torch + +from rl_engine.kernels.ops.pytorch.linear.matmul import NativeMatmulOp + +QWEN3_HIDDEN = 4096 +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, + 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 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() + 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]) + + 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( + "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)