pymc-devs / pymc-devs/pytensor

Rewrite `expand_dims` implied in vector indices

Open
#1,138 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

graph rewriting indexing numba
Dominant language
Python
Stars
644
Forks
208
Avg merge
2d 14h
Merged PRs (30d)
16

Description

Description

The model behind #1132 cannot run in non-obj numba due to an implicit expand_dims in the vector indexing, that looks like:

import pytensor
import pytensor.tensor as pt
import numpy as np

a = pt.tensor(shape=(None, None))
b = a[pt.arange(10)[:, None], pt.arange(10)[:, None]]
c = a[pt.arange(10), pt.arange(10)][:, None]

# Issues UserWarning: Numba will use object mode to run AdvancedSubtensor's perform method
fn_b = pytensor.function([a], b, mode="NUMBA")

# Runs in non-obj mode
fn_c = pytensor.function([a], c, mode="NUMBA")

test_a = pt.random.normal(size=(10, 10)).eval()
np.testing.assert_allclose(fn_b(test_a), fn_c(test_a))

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 with the vector-indexing reproducer in the issue and the AdvancedSubtensor path it exercises. Compare the implicit expand_dims case with the explicit trailing-axis case, then run the shown NUMBA compilation and equivalence checks. Done means the vector-indexing form compiles in non-object Numba mode and matches the explicit form.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
compilers
Issue type
Refactor
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.