pymc-devs / pymc-devs/pytensor

Lift Elemwise through SetSubtensor

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

Nobody has claimed this yet.

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

Description

Description

Pattern

For a unary Elemwise f and any SetSubtensor variant (basic, advanced, inc):

# Before
f(x[idx].set(c))

# After
f(x)[idx].set(f(c))

Motivating case: LKJCorr logp

The LKJCorr Cholesky factor L is built by filling a zero matrix:

L = zeros(n, n)
L = L[tril_indices].set(params)   # fill lower triangle
L = L[diag_indices].set(1)        # set diagonal to 1

The logp computes Sqr(L). Currently Sqr operates on the full (n, n) matrix. Lifting it through the SetSubtensors gives:

 └─ AdvancedSetSubtensor [diag=Sqr(1)=1]
    └─ AdvancedSetSubtensor [tril=Sqr(params)]
       └─ Alloc(Sqr(0)=0, n, n)

Now Sqr only applies to the params vector (length n*(n-1)/2) instead of the full matrix. The scalar applications Sqr(0) and Sqr(1) stay scalar / constant-fold.

When to apply

The rewrite is always valid for set mode, but only profitable when the base or the set-value is scalar/broadcast. In that case f on that piece stays a scalar op — cheap. If both base and set-value are full-sized arrays, lifting splits one f into two full-sized ones with no benefit.

Guard: apply when all(base.type.broadcastable) OR all(set_value.type.broadcastable).

Set mode only — for inc mode, f(x + inc_at_idx) ≠ f(x)[idx].set(f(inc)) in general. As always inc on zeros is semantically the same as set if there are no repeats (shouldn't we canonicalize into set then?).

Scope

  • All SetSubtensor variants: SetSubtensor, AdvancedSetSubtensor, AdvancedSetSubtensor1
  • Unary Elemwise always qualifies. Multi-input Elemwise qualifies when the other inputs are scalar/broadcast (e.g., Add(x[idx].set(c), 1)Add(x, 1)[idx].set(Add(c, 1))). If another input has full-sized dimensions, lifting can't be done without splitting the other inputs, which isn't trivial

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 locating the optimizer entry points for Elemwise and the SetSubtensor, AdvancedSetSubtensor, and AdvancedSetSubtensor1 variants. Trace the LKJCorr logp motivating case and verify the rewrite is limited to set mode, uses the broadcastability guard, and handles only supported Elemwise inputs. Done means eligible expressions are lifted while inc mode and non-beneficial full-sized cases remain unchanged.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
performance
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.