pymc-devs / pymc-devs/pytensor
MLX linker: Composite op fails with mixed-shape inputs
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 644
- Forks
- 208
- Avg merge
- 2d 14h
- Merged PRs (30d)
- 16
Description
Problem
The MLX linker's Composite dispatch fails at runtime when inputs have heterogeneous shapes (scalar, vector, matrix mixed together in one fused op).
ValueError: [stack] All arrays must have the same shape
MRE
import numpy as np
import pytensor.tensor as pt
from pytensor.compile import Mode, function
x = pt.matrix("x")
y = pt.vector("y")
z = pt.dscalar("z")
# Build expression with mixed-shape sub-expressions
expr = (x * y).sum(axis=-1) + z + (x * y).prod(axis=1) + pt.sum(x, axis=0)
fn = function([x, y, z], expr, mode=Mode(linker="mlx", optimizer="fast_run"))
fn(np.ones((26, 3)), np.ones(3), 1.0)
# ValueError: [stack] All arrays must have the same shape
Context
This prevents the pymc-marketing MMM model from using the MLX linker for MCMC sampling. The gradient graph (dlogp) produces Composite ops with many inputs of varying shapes — the Numba and CVM linkers handle this correctly.
uv pip install mlx (Apple Silicon only) to reproduce.
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 reproducing the mixed-shape expression from the MRE with the MLX linker and inspect the MLX linker's Composite dispatch. Compare its behavior with the Numba and CVM linkers, which already handle the graph, and add coverage for scalar, vector, and matrix inputs. Done means the MRE executes without the stack shape error.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- numpy, python
- Domain
- backend, machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Quiet
- Clarity
- Mostly clear
- Newbie friendliness
- 55/100