Skip to content

feat(ws1): Add PyTorch matmul reference operator - #168

Merged
frank-2077 merged 7 commits into
RL-Align:mainfrom
frank-2077:issue-108-matmul
Jul 4, 2026
Merged

frank-2077 merged 7 commits into
RL-Align:mainfrom
frank-2077:issue-108-matmul

Conversation

@frank-2077

@frank-2077 frank-2077 commented Jun 21, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Adds the PyTorch reference GEMM/Matmul operator for Issue #108.

This implements NativeMatmulOp as the fp32 ground-truth baseline for dense projection
matmuls. The operator follows the frozen #108 interface contract:

  • forward_fp32(a, b) casts inputs to fp32 and calls torch.matmul once
  • op_class = "reduction"
  • registry dispatch via kernel_registry.get_op("matmul")

Also adds operator documentation under docs/operators/matmul.md.

Implementation

  • Added rl_engine/kernels/ops/pytorch/linear/matmul.py
  • Added rl_engine/kernels/ops/pytorch/linear/__init__.py
  • Registered PYTORCH_NATIVE_MATMUL in rl_engine/kernels/registry.py
  • Added tests/test_matmul.py
  • Added Matmul operator docs and linked them from docs/operators/README.md
image

Summary by CodeRabbit

  • New Features

    • Added a native PyTorch reference matmul operator.
    • Enhanced operator dispatch so the matmul implementation is selected consistently across CPU, CUDA, and ROCm.
  • Documentation

    • Added a dedicated documentation page for the matmul operator.
    • Updated the operators index to include matmul.
  • Tests

    • Added comprehensive matmul test coverage for correctness, dtype/FP32 behavior, supported shape variants, batch invariance, backward/gradient checks, and registry integration.

@coderabbitai

coderabbitai Bot commented Jun 21, 2026 •

Copy link
Copy Markdown

Review Change Stack

Caution

Review failed

The pull request is closed.

ℹ️ Recent review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro

Run ID: 12ba9f1c-f620-4483-b5a3-0bbfbc36d97a

📥 Commits

Reviewing files that changed from the base of the PR and between 2a44ed7 and 8372f08.

📒 Files selected for processing (2)
  • docs/operators/README.md
  • rl_engine/kernels/registry.py

📝 Walkthrough

Walkthrough

Adds NativeMatmulOp, a PyTorch fp32-reference matmul operator, and wires it into kernel registry dispatch for CUDA, ROCm, and CPU. The PR also adds tests and documentation for the new operator.

Changes

NativeMatmulOp Feature

Layer / File(s) Summary
NativeMatmulOp implementation and registry wiring
rl_engine/kernels/ops/pytorch/linear/matmul.py, rl_engine/kernels/ops/pytorch/linear/__init__.py, rl_engine/kernels/registry.py
NativeMatmulOp implements forward_fp32 with a single torch.matmul on fp32-cast inputs and forward that casts the result back to the input dtype. The linear package exports NativeMatmulOp, and the registry adds OpBackend.PYTORCH_NATIVE_MATMUL plus "matmul" dispatch entries for CUDA, ROCm, and CPU.
NativeMatmulOp test suite
tests/test_matmul.py
Test cases cover output shape, dtype behavior, fp32 equivalence, input non-mutation, op_class metadata, batch invariance, gradient comparisons, dtype-parameterized accuracy, Qwen3 projection shapes, and registry lookup for "matmul".
Matmul operator documentation
docs/operators/matmul.md, docs/operators/README.md
Adds a new matmul operator page describing entry points, backend dispatch, tensor contract, Qwen3 shapes, reference semantics, accuracy expectations, test command, and related files, and links it from the operators README.

Estimated code review effort: 3 (Moderate) | ~25 minutes

Suggested reviewers: inaniloquentee, Flink-ddd, KJLdefeated

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 0.00% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly and concisely summarizes the main change: adding a PyTorch matmul reference operator.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests

Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

@frank-2077 frank-2077 changed the title Add PyTorch matmul reference operator feat(ws1): Add PyTorch matmul reference operator Jun 22, 2026

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Here are some review comments, everything else is fine.

Comment thread tests/test_matmul.py
def __init__(self) -> None:
pass

def __call__(self, a: Tensor, b: Tensor) -> Tensor:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

As with the previous WS1 reference operators, NativeMatmulOp must inherit from torch.nn.Module to ensure upstream compatibility with Dynamo tracing and PyTorch hooks.

Please inherit from nn.Module, initialize super().init(), and remove the manually defined call method.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Updated NativeMatmulOp to inherit torch.nn.Module and removed the manual call override.

Comment thread tests/test_matmul.py
@Flink-ddd

Copy link
Copy Markdown
Collaborator

The Same, Please resolve the code conflicts first. If the review information has already resolved the issue or only requires explanation, please mark it. If there are no other issues, we will merge this PR. Thank you.

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

♻️ Duplicate comments (1)
tests/test_matmul.py (1)

106-125: 🩺 Stability & Availability | 🟡 Minor | ⚡ Quick win

Forward batch-invariance tests lack the single-thread guard.

_single_threaded_torch was added (per prior review feedback) but only test_batch_grad_invariance (line 133) uses it. test_batch1_vs_batchN_bitwise and test_batch_invariance_with_padding do bitwise torch.equal comparisons on plain forward passes without pinning torch.set_num_threads(1), so on a multi-threaded CPU BLAS backend the reduction order (and thus bitwise result) is not guaranteed to be reproducible between the full-batch and single-row calls.

🔧 Proposed fix
     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}"
+        with _single_threaded_torch():
+            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)
+        with _single_threaded_torch():
+            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])

Per prior review, this is the same batch-invariance-determinism concern; only partially followed through.

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@tests/test_matmul.py` around lines 106 - 125, The forward batch-invariance
tests are doing bitwise `torch.equal` comparisons without enforcing
single-threaded execution, so their results can vary on multi-threaded CPU
backends. Update `test_batch1_vs_batchN_bitwise` and
`test_batch_invariance_with_padding` to run under the existing
`_single_threaded_torch` guard, matching the pattern already used in
`test_batch_grad_invariance`, so the full-batch and sliced calls in
`NativeMatmulOp.forward_fp32` are compared deterministically.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Duplicate comments:
In `@tests/test_matmul.py`:
- Around line 106-125: The forward batch-invariance tests are doing bitwise
`torch.equal` comparisons without enforcing single-threaded execution, so their
results can vary on multi-threaded CPU backends. Update
`test_batch1_vs_batchN_bitwise` and `test_batch_invariance_with_padding` to run
under the existing `_single_threaded_torch` guard, matching the pattern already
used in `test_batch_grad_invariance`, so the full-batch and sliced calls in
`NativeMatmulOp.forward_fp32` are compared deterministically.

ℹ️ Review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro

Run ID: 65e46bbb-2fa0-4903-8b77-c3f65b3cbd66

📥 Commits

Reviewing files that changed from the base of the PR and between c545946 and 2a44ed7.

📒 Files selected for processing (6)
  • docs/operators/README.md
  • docs/operators/matmul.md
  • rl_engine/kernels/ops/pytorch/linear/__init__.py
  • rl_engine/kernels/ops/pytorch/linear/matmul.py
  • rl_engine/kernels/registry.py
  • tests/test_matmul.py
✅ Files skipped from review due to trivial changes (2)
  • docs/operators/README.md
  • docs/operators/matmul.md
🚧 Files skipped from review as they are similar to previous changes (3)
  • rl_engine/kernels/ops/pytorch/linear/matmul.py
  • rl_engine/kernels/registry.py
  • rl_engine/kernels/ops/pytorch/linear/init.py

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM now, Thank you for update. cc @KJLdefeated PTAL again.

@KJLdefeated
KJLdefeated self-requested a review July 4, 2026 11:25

@KJLdefeated KJLdefeated left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LTGM, thx!

@frank-2077
frank-2077 merged commit c3ac59b into RL-Align:main Jul 4, 2026
3 of 4 checks passed
@frank-2077
frank-2077 deleted the issue-108-matmul branch August 12, 2026 09:12
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants