pymc-devs / pymc-devs/pytensor
Add MLX GEMM dispatch
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 644
- Forks
- 208
- Avg merge
- 2d 14h
- Merged PRs (30d)
- 16
Description
Description
mlx has a GEMM function that they call mlx.core.addmm. We can dispatch our GEMM Op to it as follows:
@mlx_funcify.register(Gemm)
def mlx_funcify_Gemm(op, **kwargs):
# GEMM has signature:
# b * z + a * dot(x, y)
def gemm(z, a, x, y, b):
# mx.addmm has signature:
# alpha * (a @ b) + beta * c
return mx.addmm(z, x, y, alpha=a, beta=b)
return gemm
What's tricky is that the blas rewrite machinery is quite C specific. I'm not sure if we can register just the dot22_to_gemm rewrite for sure in MLX.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reading the existing mlx_funcify registrations and the Gemm entry point shown in the issue, then trace the dot22_to_gemm rewrite and its BLAS-specific assumptions. Confirm how MLX exposes addmm and whether that rewrite can be registered independently; done means GEMM dispatches through addmm and the applicable rewrite path works without C-specific machinery.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- backend, compilers
- Issue type
- Feature
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Quiet
- Clarity
- Mostly clear
- Newbie friendliness
- 48/100