pymc-devs / pymc-devs/pytensor

Canonicalize associative Op argument order and add distributive factorization for Add/Sub

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

Nobody has claimed this yet.

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

Description

Description

PyTensor currently doesn't simplify a*b + a*c to a*(b + c), nor does it cancel a*b - b*a to 0. Both are missing for the same underlying reason: Mul arguments aren't canonicalized into a deterministic order, and there's no distributive factorization rewrite across Add/Sub. Concrete cases that all currently miss: a*b ± a*c, b*a ± c*a, mixed positions like a*b + c*a, the n-ary form a*b + a*c + a*d, and the unit-factor degenerate cases a ± a*c. The fix is two small changes:

(1) sort Mul arguments by a stable key during canonicalize, so a*b and b*a become the same node — this alone makes a*b - b*a → 0 fall out via the existing x - x → 0 rewrite;

(2) add a rewrite that recognizes Add(Mul(a, X_1), Mul(a, X_2), ...) and rewrites to Mul(a, Add(X_1, X_2, ...)). The same canonical-order principle applies to other commutative variadic ops (Add, Maximum, Minimum, bitwise/logical And/Or/Xor, and the axis tuple of commutative reductions); Mul is the immediate case here but the canonicalization step is worth doing uniformly.

The direct saving is one elementwise multiply per factored group: a*b + a*c does two multiplies and one add, a*(b + c) does one multiply and one add. For pure-scalar/pure-elementwise graphs Composite fusion compiles both forms to one kernel with the same per-element work, so the rewrite is mostly a no-op there.

The analogous rewrite for Dot (A@B + A@C → A@(B+C)) saves an entire matmul and is a separate, higher-priority follow-up; it would be natural to implement once the Mul version exists as a template.

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 at the canonicalize entry points for Mul and the other commutative variadic operations, then trace the Add/Sub rewrite machinery. Verify deterministic operand ordering and factorization for the listed binary, mixed-position, n-ary, and unit-factor cases; the Dot rewrite is explicitly a separate follow-up.

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.