pymc-devs / pymc-devs/pytensor

MLX linker: Composite op fails with mixed-shape inputs

Open
#2,316 1 comment 0 reactions 0 assignees View on GitHub

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

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 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.