Repository navigation
feat(ws1): Add PyTorch matmul reference operator - #168
Conversation
|
Caution Review failedThe pull request is closed. ℹ️ Recent review info⚙️ Run configurationConfiguration used: defaults Review profile: CHILL Plan: Pro Run ID: 📒 Files selected for processing (2)
📝 WalkthroughWalkthroughAdds ChangesNativeMatmulOp Feature
Estimated code review effort: 3 (Moderate) | ~25 minutes Suggested reviewers: 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
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. Comment |
Flink-ddd
left a comment
There was a problem hiding this comment.
Here are some review comments, everything else is fine.
| def __init__(self) -> None: | ||
| pass | ||
|
|
||
| def __call__(self, a: Tensor, b: Tensor) -> Tensor: |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
Updated NativeMatmulOp to inherit torch.nn.Module and removed the manual call override.
|
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. |
c545946 to
2a44ed7
Compare
There was a problem hiding this comment.
♻️ Duplicate comments (1)
tests/test_matmul.py (1)
106-125: 🩺 Stability & Availability | 🟡 Minor | ⚡ Quick winForward batch-invariance tests lack the single-thread guard.
_single_threaded_torchwas added (per prior review feedback) but onlytest_batch_grad_invariance(line 133) uses it.test_batch1_vs_batchN_bitwiseandtest_batch_invariance_with_paddingdo bitwisetorch.equalcomparisons on plain forward passes without pinningtorch.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
📒 Files selected for processing (6)
docs/operators/README.mddocs/operators/matmul.mdrl_engine/kernels/ops/pytorch/linear/__init__.pyrl_engine/kernels/ops/pytorch/linear/matmul.pyrl_engine/kernels/registry.pytests/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
left a comment
There was a problem hiding this comment.
LGTM now, Thank you for update. cc @KJLdefeated PTAL again.
Summary
Adds the PyTorch reference GEMM/Matmul operator for Issue #108.
This implements
NativeMatmulOpas the fp32 ground-truth baseline for dense projectionmatmuls. The operator follows the frozen #108 interface contract:
forward_fp32(a, b)casts inputs to fp32 and callstorch.matmulonceop_class = "reduction"kernel_registry.get_op("matmul")Also adds operator documentation under
docs/operators/matmul.md.Implementation
rl_engine/kernels/ops/pytorch/linear/matmul.pyrl_engine/kernels/ops/pytorch/linear/__init__.pyPYTORCH_NATIVE_MATMULinrl_engine/kernels/registry.pytests/test_matmul.pydocs/operators/README.mdSummary by CodeRabbit
New Features
matmuloperator.matmulimplementation is selected consistently across CPU, CUDA, and ROCm.Documentation
matmuloperator.matmul.Tests
matmultest coverage for correctness, dtype/FP32 behavior, supported shape variants, batch invariance, backward/gradient checks, and registry integration.