Repository navigation
feat(ws1): Add PyTorch matmul reference operator #168
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
7 commits
Select commit
Hold shift + click to select a range
29e360b
Add PyTorch matmul reference operator
frank-2077 f364d8f
format with pre-commit
frank-2077 64bd581
Make native matmul a torch module
frank-2077 65e2c23
Add native matmul backward test
frank-2077 2a44ed7
Add native matmul batch grad invariance test
frank-2077 3746f0b
Merge branch 'main' into issue-108-matmul
frank-2077 8372f08
Merge branch 'main' into issue-108-matmul
frank-2077 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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` |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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"] |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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()) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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}" | ||
| ) | ||
|
|
||
|
frank-2077 marked this conversation as resolved.
|
||
|
|
||
| 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) | ||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.