pymc-devs / pymc-devs/pytensor

Add MLX GEMM dispatch

Open
#2,008 7 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

graph rewriting linalg mlx Op implementation
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

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.