Skip to content

feat(ir): linear-algebra simplification rewrites [#512] - #1016

Open
Cojo01 wants to merge 74 commits into
daphne-project:mainfrom
Cojo01:512-linalg-rewrites
Open

Cojo01 wants to merge 74 commits into
daphne-project:mainfrom
Cojo01:512-linalg-rewrites

Conversation

@Cojo01

@Cojo01 Cojo01 commented Jul 10, 2026 •

Copy link
Copy Markdown

Fourth of four PRs for #512. It adds a set of linear-algebra rewrites, building on the earlier three PRs.

The four PRs:

  1. refactor(ir): Mark pure DaphneIR base classes as Pure [#512] #1013; purity traits, plus retiring the CSE workaround.
  2. feat(ir): Algebraic trait vocabulary + shared scalar-fold interface [#512] #1014; the ScalarConstantFoldable interface and the algebraic-trait vocabulary; moving existing rewrites onto it and adding new trait-based rewrites.
  3. feat(ir): Fail-closed property accessors, and make symmetric a real inferred property [#512] #1015; fail-closed property accessors, symmetric inference for matrix products, and the transpose elision it unlocks.
  4. feat(ir): linear-algebra simplification rewrites [#512] #1016 (this PR); the linear-algebra rewrites below.

The rewrites

All hand-written OpRewritePatterns in a new src/ir/daphneir/LinearAlgebraRewrites.{cpp,h}, registered into the existing --daphne-algebraic-simplify pass. No new pass, CLI flag, or Passes.td change.

  • Trace of a product sum(diagVector(X @ Y)) → sum(X * t(Y)). DAPHNE has no trace builtin, so users write the diagonal-of-a-product form, which builds the full n x n product just to keep its diagonal and drop the rest. The element-wise form gives the same sum without ever building the product.
  • Row scaling diag(v) @ X → X * v, and column scaling X @ diag(v) → X * t(v). Multiplying by a diagonal matrix just scales each row (or column), which an element-wise multiply does directly; so the big diagonal matrix and the matmul both go away.
  • Pulling a scalar out of a sum sum(s * X) → s * sum(X). Multiply once after the sum instead of once per element before it.
  • Dropping a length-1 aggregate rowAgg(X) / colAgg(X) when the axis being reduced has length 1. There is nothing to combine, so the aggregate returns its input and is removed. Covers the sum / min / max row and column variants.
  • Turning repeated adds into a multiply X + X → X * 2, plus (X * c) + X → X * (c + 1) and (X * c1) + (X * c2) → X * (c1 + c2). These chain, so X + X + X becomes X * 3.
  • Combining two shifted matrices (M1 + s1) + (M2 + s2) → (M1 + M2) + (s1 + s2). Add the two scalars once instead of spreading each one across a whole matrix.

The change also marks DiagVectorOp Pure (its sibling DiagMatrixOp already was) so dead-code elimination can remove the leftover matMul / diagVector after the trace rewrite fires. The other rewrites already relied on TransposeOp / EwMulOp being Pure.

Soundness notes

Each rewrite does nothing unless its guards hold. They read shapes through the fail-closed accessors from #1015 and give up on anything unknown. The ones worth calling out:

  • Result type of the sum. sum(s * X) adds up in the product's type, s * sum(X) in X's type. If s is wider than X (say f64 times si64), the rewritten integer sum can overflow where the original would not, so this rewrite only fires when X's type already matches the sum's result type.
  • Floating point. The regrouping rewrites are integer-only, because reordering additions can change floating-point rounding. X + X → X * 2 is the one exception: doubling a float is exact, so it stays bit-for-bit identical (NaN / Inf / -0.0 included). Integer overflow in the regrouped scalar sums is fine; DAPHNE's integer kernels wrap at 2^n, so the result is identical even when an intermediate wraps.
  • Types on new ops. Inference runs before this pass, so new ops are given real result types from the known sizes rather than an unknown type that would never get resolved later.
  • Operand order (scaling). The element-wise kernel only broadcasts when the matrix is the left operand, so the rewrites put the matrix on the left and the vector on the right. The EwMul canonicalizer only swaps a scalar left operand, so that order sticks.

Correctness

Each rewrite has both a LIT test (test/codegen/rewrite_linalg_simplify.mlir: a positive case, plus the negatives that must not fire: unknown shapes, transposed matmul flags, and the type cases above) and an end-to-end numerical test (test/api/cli/operations/) that runs a triggering script and an equivalent reference script written a way the rewrite can't match, and checks the output matches exactly. The repeated-add and shifted-matrix pairs deliberately overflow si64 / si32, so identical output confirms the wrapping arithmetic is preserved.

Those numerical tests paid off. The first version of the trace rewrite produced an unknown-typed product that compiled on its own but failed to lower on a real script, because the type was never resolved and kernel dispatch had nothing to bind to. A test that only checks whether the rewrite fires would have missed it. The fix, giving the new ops a real result type, is included here.

Full [operations] suite green (488 assertions / 67 cases). Otherwise the suite matches baseline. The failures the local wrapper filters out are two pre-existing aarch64-only crashes (#1011, #1012) that don't reproduce on x86_64 CI.

Deliberately out of scope

  • Weighted operators (wsloss/wsigmoid/wdivmm/wcemm); no ops or kernels for them in DAPHNE yet, so they are follow-up work.
  • Three rewrites from 512 Simplification Rewrites for the Open Source Project Daphne #974, each left out for a specific reason. Slice-of-matmul: the canonicalizer that runs before this pass already covers the static cases, and a fail-closed version can't reach the dynamic ones. Full-cover insert elision: unsound as it would have to be written, because ShapeFromArg makes the column check always pass, so it would silently drop partial writes. Zero-absorption based on sparsity: a compile-time sparsity estimate isn't a guarantee about the actual values. The sum(t(X)) / sum(reverse(X)) fold an earlier version carried now lives in feat(ir): Algebraic trait vocabulary + shared scalar-fold interface [#512] #1014, via the OnlyReordersElements / OrderAgnosticAggregate traits.

On the commit range

This sits on top of #1015. The base branches live on my fork and can't be a PR base upstream, so the diff currently also includes the commits from #1013, #1014, and #1015. Once those merge I'll rebase and force-push, leaving only this PR's own commits.

Refs #512.

@Cojo01
Cojo01 force-pushed the 512-linalg-rewrites branch 3 times, most recently from 8f538af to bd752de Compare July 12, 2026 21:42
@Cojo01
Cojo01 force-pushed the 512-linalg-rewrites branch from db5677a to cc78404 Compare July 14, 2026 12:58
@Cojo01 Cojo01 changed the title feat(ir): linear-algebra simplification rewrites (trace idiom, diagonal scaling) feat(ir): linear-algebra simplification rewrites [#512] Jul 18, 2026
@Cojo01
Cojo01 force-pushed the 512-linalg-rewrites branch 2 times, most recently from a1fd6b8 to cb20bd6 Compare July 18, 2026 21:19
@Cojo01
Cojo01 force-pushed the 512-linalg-rewrites branch from 087fbbe to 4d5786d Compare July 19, 2026 12:24
@Cojo01
Cojo01 force-pushed the 512-linalg-rewrites branch from 4d5786d to e857aee Compare July 19, 2026 13:14
@Cojo01
Cojo01 marked this pull request as ready for review July 19, 2026 16:23

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants