From a17aee58b2ce0d0067ed4504faa1a61a57fa89a0 Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Sun, 12 Jul 2026 02:13:05 -0400
Subject: [PATCH 1/9] Fix BufferizationStage IR amplification for
runtime-coefficient Hamiltonians
Add scalarize-tensor-extracts and reroll-loops passes and re-order the
default pipeline so runtime-coefficient Trotterization no longer amplifies
IR through bufferization (9.9x -> ~1.0x on the H2 QPE benchmark).
Fixes #2759
---
doc/releases/changelog-dev.md | 17 +
frontend/test/lit/test_detensorize.py | 8 +-
.../pytest/test_trotter_runtime_coeffs.py | 114 ++++
mlir/include/Catalyst/Transforms/Passes.td | 64 +++
.../DefaultPipelines/DefaultPipelines.h | 24 +-
mlir/lib/Catalyst/Transforms/CMakeLists.txt | 2 +
.../Catalyst/Transforms/RerollLoopsPass.cpp | 518 ++++++++++++++++++
.../ScalarizeTensorExtractsPass.cpp | 283 ++++++++++
mlir/test/Catalyst/RerollLoopsTest.mlir | 133 +++++
.../Catalyst/ScalarizeTensorExtractsTest.mlir | 146 +++++
10 files changed, 1302 insertions(+), 7 deletions(-)
create mode 100644 frontend/test/pytest/test_trotter_runtime_coeffs.py
create mode 100644 mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
create mode 100644 mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
create mode 100644 mlir/test/Catalyst/RerollLoopsTest.mlir
create mode 100644 mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index 771149a913..c0b532c3c4 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -12,6 +12,23 @@
Improvements ðŸ›
+* IR amplification for circuits whose gate parameters are computed from runtime values
+ (e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
+ Three changes work together in the default pipeline:
+ the new `scalarize-tensor-extracts` pass sinks scalar `tensor.extract` operations through
+ `tensor.extract_slice`, `tensor.collapse_shape`, and small elementwise `linalg.generic`
+ producers so gate-angle dataflow becomes pure scalar arithmetic instead of thousands of
+ tiny tensors that each survive bufferization as an allocation and copies;
+ `linalg-fuse-elementwise-ops` plus additional `cse` applications (before and after
+ `one-shot-bufferize`) remove the duplicate angle computations produced by tracing; and
+ the new `reroll-loops` pass reconstructs the loops that Python tracing unrolled, by
+ detecting tandem repeats of structurally isomorphic operation windows (via structural
+ hashing), verifying that cross-window dataflow is limited to threaded SSA values (such as
+ qubit values) plus loop-invariant values, and replacing each repeat with an `scf.for`.
+ For the H2 QPE benchmark from the issue, gate-op volume after HLO lowering drops about
+ 5x and downstream IR, compile time, and peak memory drop accordingly.
+ [(#2759)](https://github.com/PennyLaneAI/catalyst/issues/2759)
+
* The `decompose-lowering` pass now supports applying a selection of the available decomposition rules via the `target_rules` parameter.
The pass also no longer applies the `inline`, `cse` and `canonicalize` passes to avoid unnecessary IR mutations.
Instead, decomposition rules are deterministically inlined by a custom function (`inline` is non-deterministic, using an estimated benefit and threshold as criteria for inlining).
diff --git a/frontend/test/lit/test_detensorize.py b/frontend/test/lit/test_detensorize.py
index cc2a2295cd..40cf98509b 100644
--- a/frontend/test/lit/test_detensorize.py
+++ b/frontend/test/lit/test_detensorize.py
@@ -45,15 +45,15 @@ def f_with_cond(a, b):
# CHECK-NOT: linalg.generic
a2 = a + a
if a2 > b:
- # CHECK: arith.subf
+ # CHECK-DAG: arith.subf
a = a - 2.0
- # CHECK: arith.mulf
+ # CHECK-DAG: arith.mulf
b = b * 2.0
c = a + b
else:
- # CHECK: arith.mulf
+ # CHECK-DAG: arith.mulf
a = a * 2.0
- # CHECK: arith.subf
+ # CHECK-DAG: arith.subf
b = b - 2.0
c = a + b
return c * 2.0
diff --git a/frontend/test/pytest/test_trotter_runtime_coeffs.py b/frontend/test/pytest/test_trotter_runtime_coeffs.py
new file mode 100644
index 0000000000..15f446a5dd
--- /dev/null
+++ b/frontend/test/pytest/test_trotter_runtime_coeffs.py
@@ -0,0 +1,114 @@
+# Copyright 2026 Xanadu Quantum Technologies Inc.
+
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+
+# http://www.apache.org/licenses/LICENSE-2.0
+
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Integration tests for the IR-amplification fixes for runtime-coefficient
+Hamiltonians (issue #2759): scalarize-tensor-extracts, elementwise fusion, and
+reroll-loops in the default pipeline must preserve numerics for Trotterized
+workloads with runtime coefficients."""
+
+import numpy as np
+import pennylane as qml
+import pytest
+from jax import numpy as jnp
+
+from catalyst import qjit
+
+# H2/STO-3G-like coefficients; the structure (15 terms, runtime values) is what
+# exercises the pipeline, the values just need to be a valid Hamiltonian.
+COEFFS = [
+ -0.0996, 0.1711, 0.1711, -0.2225, -0.2225, 0.1686, 0.0453, -0.0453,
+ -0.0453, 0.0453, 0.1205, 0.1658, 0.1658, 0.1205, 0.1743,
+]
+OPS_FACTORY = lambda: [
+ qml.Identity(0),
+ qml.PauliZ(0), qml.PauliZ(1), qml.PauliZ(2), qml.PauliZ(3),
+ qml.PauliZ(0) @ qml.PauliZ(1),
+ qml.PauliY(0) @ qml.PauliX(1) @ qml.PauliX(2) @ qml.PauliY(3),
+ qml.PauliY(0) @ qml.PauliY(1) @ qml.PauliX(2) @ qml.PauliX(3),
+ qml.PauliX(0) @ qml.PauliX(1) @ qml.PauliY(2) @ qml.PauliY(3),
+ qml.PauliX(0) @ qml.PauliY(1) @ qml.PauliY(2) @ qml.PauliX(3),
+ qml.PauliZ(0) @ qml.PauliZ(2), qml.PauliZ(0) @ qml.PauliZ(3),
+ qml.PauliZ(1) @ qml.PauliZ(2), qml.PauliZ(1) @ qml.PauliZ(3),
+ qml.PauliZ(2) @ qml.PauliZ(3),
+]
+
+N_QUBITS = 4
+N_EST = 2
+N_TROTTER = 6 # enough repetitions for reroll-loops to fire
+
+
+def make_qpe(dev, runtime: bool):
+ """Controlled-power Trotterized QPE, with runtime or compile-time coeffs."""
+
+ @qml.qnode(dev)
+ def qpe_circuit(coeffs):
+ qml.PauliX(0)
+ qml.PauliX(1)
+ for k in range(N_EST):
+ qml.Hadamard(wires=N_QUBITS + k)
+ H = qml.dot(coeffs if runtime else COEFFS, OPS_FACTORY())
+ for k in range(N_EST):
+ t = 2 ** (N_EST - 1 - k)
+ qml.ctrl(
+ qml.adjoint(
+ qml.TrotterProduct(H, time=t, n=N_TROTTER, order=2,
+ check_hermitian=False)
+ ),
+ control=N_QUBITS + k,
+ )
+ qml.adjoint(qml.QFT)(wires=range(N_QUBITS, N_QUBITS + N_EST))
+ return qml.probs(wires=range(N_QUBITS, N_QUBITS + N_EST))
+
+ return qpe_circuit
+
+
+class TestRuntimeCoefficientTrotter:
+ """Numerical equivalence of runtime- and fixed-coefficient Trotterization
+ through the default pipeline (which scalarizes, fuses, and rerolls)."""
+
+ def test_runtime_matches_fixed(self):
+ """qml.dot with traced coefficients must produce the same distribution
+ as the same Hamiltonian with Python-float coefficients."""
+ dev = qml.device("lightning.qubit", wires=N_QUBITS + N_EST)
+ coeffs = jnp.array(COEFFS)
+
+ dyn = qjit(make_qpe(dev, runtime=True))(coeffs)
+ fixed = qjit(make_qpe(dev, runtime=False))(coeffs)
+
+ assert np.allclose(np.asarray(dyn), np.asarray(fixed), atol=1e-9)
+
+ def test_reroll_recovers_loops(self):
+ """The compiled module must contain scf.for loops recovered from the
+ unrolled Trotter steps (guards against silent regression of
+ reroll-loops in the default pipeline)."""
+ dev = qml.device("lightning.qubit", wires=N_QUBITS + N_EST)
+ coeffs = jnp.array(COEFFS)
+
+ compiled = qjit(make_qpe(dev, runtime=True), keep_intermediate=True)
+ compiled(coeffs)
+ try:
+ workspace = str(compiled.workspace)
+ import glob
+ import os
+
+ hlo_files = glob.glob(os.path.join(workspace, "*HLOLowering*.mlir"))
+ assert hlo_files, "no post-HLO snapshot written"
+ content = open(hlo_files[0], encoding="utf-8").read()
+ assert "scf.for" in content, "reroll-loops did not fire"
+ finally:
+ compiled.workspace.cleanup()
+
+
+if __name__ == "__main__":
+ pytest.main(["-x", __file__])
diff --git a/mlir/include/Catalyst/Transforms/Passes.td b/mlir/include/Catalyst/Transforms/Passes.td
index c5f87f0d7c..30eddd487f 100644
--- a/mlir/include/Catalyst/Transforms/Passes.td
+++ b/mlir/include/Catalyst/Transforms/Passes.td
@@ -333,6 +333,70 @@ def RegisterDecompRuleResourcePass : Pass<"register-decomp-rule-resource"> {
}];
}
+def RerollLoopsPass : Pass<"reroll-loops"> {
+ let summary = "Reconstruct loops from unrolled repetitive op sequences.";
+ let description = [{
+ Tracing a Python program unrolls its loops: N structurally identical
+ circuit segments (Trotter steps, layers, folds) arrive as N copies of the
+ same op sequence, and every downstream stage amplifies each copy (see
+ issue #2759).
+
+ This pass detects maximal tandem repeats of structurally isomorphic op
+ windows via shift-invariant structural hashing, verifies that
+ cross-window dataflow is limited to values threaded from the directly
+ preceding window (e.g. qubit SSA values) plus loop-invariant values, and
+ replaces the repeat with an `scf.for` whose iter_args are the threaded
+ values. A repeat of multiplicity k shrinks that region k-fold.
+ }];
+
+ let dependentDialects = [
+ "mlir::arith::ArithDialect",
+ "mlir::scf::SCFDialect"
+ ];
+
+ let options = [
+ Option<
+ /*C++ var name=*/"minPeriod",
+ /*CLI arg name=*/"min-period",
+ /*type=*/"unsigned",
+ /*default=*/"2",
+ /*description=*/"Minimum number of ops per repeated window."
+ >,
+ Option<
+ /*C++ var name=*/"minSavings",
+ /*CLI arg name=*/"min-savings",
+ /*type=*/"unsigned",
+ /*default=*/"8",
+ /*description=*/
+ "Minimum number of ops a reroll must eliminate to be applied."
+ >
+ ];
+}
+
+def ScalarizeTensorExtractsPass : Pass<"scalarize-tensor-extracts"> {
+ let summary = "Sink scalar tensor.extract ops through small-tensor producers.";
+ let description = [{
+ Programs that compute quantum gate parameters from runtime inputs (e.g.
+ Trotterization with runtime Hamiltonian coefficients) produce long chains of
+ small tensor operations whose only consumers are scalar `tensor.extract`
+ operations. Each intermediate tensor survives bufferization as an allocation
+ with copies and each `linalg.generic` is later unrolled into loops, causing
+ severe IR amplification (see issue #2759).
+
+ This pass folds `tensor.extract` through elementwise `linalg.generic` (by
+ inlining the scalar payload), `tensor.collapse_shape`, and
+ `tensor.extract_slice`, so the extracted element is computed directly in
+ scalar arithmetic and the intermediate tensors become dead. Payload inlining
+ is limited to small statically-shaped tensors to bound code growth.
+ }];
+
+ let dependentDialects = [
+ "mlir::arith::ArithDialect",
+ "mlir::linalg::LinalgDialect",
+ "mlir::tensor::TensorDialect"
+ ];
+}
+
def EmptyPass : Pass<"empty"> {
let summary = "Empty pass that does nothing.";
diff --git a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
index 9c652e7732..d2d91880a9 100644
--- a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
+++ b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
@@ -71,10 +71,27 @@ const PipelineList pipelineList{
"scatter-lowering",
"hlo-custom-call-lowering",
"cse",
+ // Sink scalar extractions through small-tensor producers and fuse the
+ // remaining elementwise ops. Traced gate-parameter dataflow (e.g. runtime
+ // Hamiltonian coefficients, issue #2759) otherwise reaches bufferization
+ // as thousands of tiny tensor ops that each become an alloc + copy.
+ "scalarize-tensor-extracts",
+ "func.func(linalg-fuse-elementwise-ops)",
+ "canonicalize",
+ "cse",
+ // Reconstruct the loops that tracing unrolled (e.g. Trotter steps); a
+ // repeat of multiplicity k shrinks its region k-fold before the
+ // bufferization and LLVM stages amplify it (issue #2759).
+ "reroll-loops",
"func.func(linalg-detensorize{aggressive-mode})",
"detensorize-scf",
"detensorize-function-boundary",
+ // Detensorization is what materializes tensor.extract on the gate-angle
+ // dataflow, so scalarization must run again here to fold the
+ // extract_slice/collapse_shape chains it exposes (issue #2759).
+ "scalarize-tensor-extracts",
"canonicalize",
+ "cse",
"symbol-dce"}},
{"gradient-lowering-stage",
{"annotate-invalid-gradient-functions",
@@ -93,6 +110,9 @@ const PipelineList pipelineList{
*/
// This pass is needed to avoid aliasing of the input buffer with the output buffer.
"mark-entry-point-args-non-writable",
+ // Value-number duplicate tensor computations before bufferization so they
+ // do not each become a separate buffer (issue #2759).
+ "cse",
"one-shot-bufferize",
// Remove dead memrefToTensorOp's
"canonicalize",
@@ -112,9 +132,7 @@ const PipelineList pipelineList{
// Must be after convert-bufferization-to-memref.
// Otherwise, there are issues in the lowering of dynamic tensors.
"canonicalize",
- /* [DISABLED PASS]
- * "cse",
- */
+ "cse",
"cp-global-memref"}},
{"llvm-dialect-lowering-stage",
{"qnode-to-async-lowering",
diff --git a/mlir/lib/Catalyst/Transforms/CMakeLists.txt b/mlir/lib/Catalyst/Transforms/CMakeLists.txt
index bb794d3b4c..2ecd8c019a 100644
--- a/mlir/lib/Catalyst/Transforms/CMakeLists.txt
+++ b/mlir/lib/Catalyst/Transforms/CMakeLists.txt
@@ -22,7 +22,9 @@ file(GLOB SRC
qnode_to_async_lowering.cpp
QnodeToAsyncPatterns.cpp
RegisterInactiveCallbackPass.cpp
+ RerollLoopsPass.cpp
ResourceAnalysisPass.cpp
+ ScalarizeTensorExtractsPass.cpp
RegisterDecompRuleResourcePass.cpp
SplitMultipleTapes.cpp
TBAAPatterns.cpp
diff --git a/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
new file mode 100644
index 0000000000..952188e4e8
--- /dev/null
+++ b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
@@ -0,0 +1,518 @@
+// Copyright 2026 Xanadu Quantum Technologies Inc.
+
+// Licensed under the Apache License, Version 2.0 (the "License");
+// you may not use this file except in compliance with the License.
+// You may obtain a copy of the License at
+
+// http://www.apache.org/licenses/LICENSE-2.0
+
+// Unless required by applicable law or agreed to in writing, software
+// distributed under the License is distributed on an "AS IS" BASIS,
+// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+// See the License for the specific language governing permissions and
+// limitations under the License.
+
+#define DEBUG_TYPE "reroll-loops"
+
+#include "llvm/ADT/DenseMap.h"
+#include "llvm/ADT/Hashing.h"
+#include "llvm/Support/Debug.h"
+
+#include "mlir/Dialect/Arith/IR/Arith.h"
+#include "mlir/Dialect/SCF/IR/SCF.h"
+#include "mlir/IR/Builders.h"
+#include "mlir/IR/IRMapping.h"
+#include "mlir/Pass/Pass.h"
+
+using namespace llvm;
+using namespace mlir;
+
+// Tracing a Python circuit destroys its loop structure: N structurally identical
+// Trotter steps (or layers, folds, powers) arrive as N copies of the same op
+// sequence. Every later stage (bufferization, LLVM conversion, LLVM codegen)
+// amplifies each copy, which is the dominant source of IR blowup in issue #2759.
+//
+// This pass reconstructs those loops:
+//
+// 1. Canonical structural hashing: each op in a block is hashed over its name,
+// attributes, result types, and operand provenance, where operands defined
+// inside the block are identified by their *backward distance* in ops. The
+// hash sequence is shift-invariant, so identical iterations produce
+// identical hash subsequences regardless of position.
+//
+// 2. Tandem-repeat detection: maximal runs with H[i] == H[i+p] identify
+// candidate repeats w^k of period p (a linear-time shift-and-compare over
+// candidate periods harvested from same-hash gap statistics).
+//
+// 3. Semantic verification: consecutive windows must be isomorphic op-by-op,
+// with cross-window dataflow restricted to values *threaded* from the
+// directly preceding window at a fixed (position, result) slot — exactly the
+// shape of qubit values chained through a gate sequence — and all other
+// operands loop-invariant.
+//
+// 4. Materialization: the repeat is replaced by an scf.for whose iter_args are
+// the threaded slots, its body a clone of one window. A repeat of
+// multiplicity k shrinks that region k-fold before any downstream stage
+// sees it.
+
+namespace {
+
+//===----------------------------------------------------------------------===//
+// Structural hashing
+//===----------------------------------------------------------------------===//
+
+/// Ops that may participate in a rerolled window: single-block-region-free,
+/// no successors, and not a terminator.
+bool isRerollableOp(Operation *op)
+{
+ return op->getNumRegions() == 0 && op->getNumSuccessors() == 0 &&
+ !op->hasTrait();
+}
+
+/// Compute shift-invariant structural hashes for all ops of a block (excluding
+/// the terminator). Non-rerollable ops receive unique sentinel hashes so they
+/// can never be part of a repeat.
+///
+/// In-block defs are hashed with a constant tag plus the result number:
+/// hashing them by identity would make every window distinct (each window
+/// threads different SSA values), and hashing them by backward distance breaks
+/// on loop-invariant values (e.g. CSE-deduplicated angle computations) whose
+/// distance grows window over window. The resulting weaker discrimination is
+/// compensated by semantic verification of every candidate. Block arguments
+/// and out-of-block values are hashed by identity (loop-invariant by
+/// construction).
+void computeHashes(ArrayRef ops, const llvm::DenseMap &indexOf,
+ SmallVectorImpl &hashes)
+{
+ uint64_t sentinel = 0;
+ for (Operation *op : ops) {
+ if (!isRerollableOp(op)) {
+ hashes.push_back(hash_combine(0xdeadbeefULL, ++sentinel));
+ continue;
+ }
+ hash_code h = hash_combine(op->getName().getTypeID(),
+ op->getAttrDictionary().getAsOpaquePointer());
+ for (Type t : op->getResultTypes()) {
+ h = hash_combine(h, t.getAsOpaquePointer());
+ }
+ for (Value v : op->getOperands()) {
+ Operation *def = v.getDefiningOp();
+ auto it = def ? indexOf.find(def) : indexOf.end();
+ if (it != indexOf.end()) {
+ h = hash_combine(h, 1, cast(v).getResultNumber());
+ }
+ else {
+ // Block argument or out-of-block def: hash the value identity.
+ h = hash_combine(h, 2, v.getAsOpaquePointer());
+ }
+ }
+ hashes.push_back(h);
+ }
+}
+
+//===----------------------------------------------------------------------===//
+// Candidate detection
+//===----------------------------------------------------------------------===//
+
+struct Candidate {
+ size_t start; // index of the first op of the first window
+ size_t period; // ops per window
+ size_t count; // number of windows
+};
+
+/// Harvest candidate periods from the gap distribution between successive
+/// occurrences of equal hashes, then find maximal H[i] == H[i+p] runs.
+void findCandidates(ArrayRef hashes, unsigned minPeriod, unsigned maxPeriods,
+ SmallVectorImpl &out)
+{
+ size_t n = hashes.size();
+ llvm::DenseMap lastSeen;
+ llvm::DenseMap gapWeight;
+ for (size_t i = 0; i < n; ++i) {
+ auto [it, inserted] = lastSeen.try_emplace(hashes[i], i);
+ if (!inserted) {
+ size_t gap = i - it->second;
+ if (gap >= minPeriod) {
+ gapWeight[gap]++;
+ }
+ it->second = i;
+ }
+ }
+
+ SmallVector> periods(gapWeight.begin(), gapWeight.end());
+ // Favor periods that explain the most ops.
+ llvm::sort(periods, [](auto &a, auto &b) {
+ return a.first * a.second > b.first * b.second;
+ });
+ if (periods.size() > maxPeriods) {
+ periods.resize(maxPeriods);
+ }
+
+ for (auto &[p, weight] : periods) {
+ // Require the period to repeat enough times to be worth a loop.
+ if (weight < 2 * p) {
+ continue;
+ }
+ size_t runStart = SIZE_MAX;
+ for (size_t i = 0; i + p <= n; ++i) {
+ bool match = (i + p < n) && hashes[i] == hashes[i + p];
+ if (match && runStart == SIZE_MAX) {
+ runStart = i;
+ }
+ if (!match && runStart != SIZE_MAX) {
+ size_t runLen = i - runStart;
+ size_t count = runLen / p + 1;
+ if (count >= 3) {
+ out.push_back({runStart, p, count});
+ }
+ runStart = SIZE_MAX;
+ }
+ }
+ }
+}
+
+//===----------------------------------------------------------------------===//
+// Semantic verification
+//===----------------------------------------------------------------------===//
+
+struct RerollPlan {
+ Candidate cand;
+ // Threaded slots: values flowing window -> next window, identified by
+ // (defining op position within the window, result number). Order defines
+ // the iter_args order.
+ SmallVector> slots;
+ // Initial value of each slot (operand of the first window).
+ SmallVector inits;
+ // Classification of each cross-window operand use:
+ // (window-relative op position, operand number) -> slot index.
+ llvm::DenseMap, unsigned> threadedUse;
+};
+
+/// Verify that the candidate's windows are isomorphic with dataflow limited to
+/// threaded slots + loop-invariant values, and build the reroll plan.
+std::optional verifyCandidate(ArrayRef ops,
+ const llvm::DenseMap &indexOf,
+ Candidate cand)
+{
+ size_t start = cand.start, p = cand.period, count = cand.count;
+ size_t end = start + count * p;
+ if (end > ops.size()) {
+ return std::nullopt;
+ }
+
+ RerollPlan plan;
+ plan.cand = cand;
+ llvm::DenseMap, unsigned> slotIndex; // (defPos,resNo) -> idx
+
+ auto getSlot = [&](unsigned defPos, unsigned resNo) -> unsigned {
+ auto [it, inserted] = slotIndex.try_emplace({defPos, resNo}, plan.slots.size());
+ if (inserted) {
+ plan.slots.push_back({defPos, resNo});
+ plan.inits.push_back(Value());
+ }
+ return it->second;
+ };
+
+ for (size_t w = 1; w < count; ++w) {
+ for (size_t j = 0; j < p; ++j) {
+ Operation *a = ops[start + (w - 1) * p + j];
+ Operation *b = ops[start + w * p + j];
+ if (!isRerollableOp(a) || !isRerollableOp(b)) {
+ return std::nullopt;
+ }
+ if (a->getName() != b->getName() ||
+ a->getAttrDictionary() != b->getAttrDictionary() ||
+ a->getResultTypes() != b->getResultTypes() ||
+ a->getNumOperands() != b->getNumOperands()) {
+ return std::nullopt;
+ }
+ for (unsigned t = 0; t < b->getNumOperands(); ++t) {
+ Value vb = b->getOperand(t);
+ Value va = a->getOperand(t);
+ Operation *defB = vb.getDefiningOp();
+ auto itB = defB ? indexOf.find(defB) : indexOf.end();
+ size_t ib = (itB != indexOf.end()) ? itB->second : SIZE_MAX;
+
+ if (ib != SIZE_MAX && ib >= start + w * p) {
+ // Within current window: the counterpart must reference the
+ // same relative position.
+ Operation *defA = va.getDefiningOp();
+ auto itA = defA ? indexOf.find(defA) : indexOf.end();
+ if (itA == indexOf.end() || itA->second + p != ib ||
+ cast(va).getResultNumber() !=
+ cast(vb).getResultNumber()) {
+ return std::nullopt;
+ }
+ }
+ else if (ib != SIZE_MAX && ib >= start + (w - 1) * p) {
+ // Threaded from the previous window.
+ unsigned defPos = ib - (start + (w - 1) * p);
+ unsigned resNo = cast(vb).getResultNumber();
+ unsigned slot = getSlot(defPos, resNo);
+ auto [uit, uinserted] = plan.threadedUse.try_emplace({(unsigned)j, t}, slot);
+ if (!uinserted && uit->second != slot) {
+ return std::nullopt;
+ }
+ // The counterpart operand must thread identically.
+ if (w == 1) {
+ // va is the init value; it must be loop-invariant w.r.t.
+ // the region (defined before it).
+ Operation *defA = va.getDefiningOp();
+ auto itA = defA ? indexOf.find(defA) : indexOf.end();
+ if (itA != indexOf.end() && itA->second >= start) {
+ return std::nullopt;
+ }
+ if (plan.inits[slot] && plan.inits[slot] != va) {
+ return std::nullopt;
+ }
+ plan.inits[slot] = va;
+ }
+ else {
+ Operation *defA = va.getDefiningOp();
+ auto itA = defA ? indexOf.find(defA) : indexOf.end();
+ if (itA == indexOf.end() || itA->second + p != ib ||
+ cast(va).getResultNumber() != resNo) {
+ return std::nullopt;
+ }
+ }
+ }
+ else {
+ // Loop-invariant: must be the exact same value, defined
+ // before the region.
+ if (va != vb) {
+ return std::nullopt;
+ }
+ if (ib != SIZE_MAX && ib >= start) {
+ return std::nullopt;
+ }
+ }
+ }
+ }
+ }
+
+ // Every slot must have an init.
+ for (Value init : plan.inits) {
+ if (!init) {
+ return std::nullopt;
+ }
+ }
+
+ // Uses of window results outside the allowed range:
+ // - windows 0..count-2: results may only be used inside their own window or
+ // the next one (threaded uses were verified above; any other use pattern
+ // is unsupported).
+ // - last window: external uses allowed only for threaded slots (they become
+ // loop results).
+ for (size_t w = 0; w < count; ++w) {
+ bool isLast = (w == count - 1);
+ for (size_t j = 0; j < p; ++j) {
+ Operation *op = ops[start + w * p + j];
+ for (OpResult res : op->getResults()) {
+ for (OpOperand &use : res.getUses()) {
+ Operation *user = use.getOwner();
+ auto uit = indexOf.find(user);
+ size_t ui = (uit != indexOf.end()) ? uit->second : SIZE_MAX;
+ bool inOwnWindow =
+ ui != SIZE_MAX && ui >= start + w * p && ui < start + (w + 1) * p;
+ bool inNextWindow = !isLast && ui != SIZE_MAX &&
+ ui >= start + (w + 1) * p &&
+ ui < start + (w + 2) * p;
+ if (inOwnWindow || inNextWindow) {
+ continue;
+ }
+ // External use.
+ if (!isLast) {
+ return std::nullopt;
+ }
+ if (!slotIndex.count({(unsigned)j, res.getResultNumber()})) {
+ return std::nullopt;
+ }
+ }
+ }
+ }
+ }
+
+ return plan;
+}
+
+/// Try to extend a verified candidate by whole windows to the left/right; the
+/// hash sequence misses the first window (its cross-window references point at
+/// the prologue, at different distances), so this recovers it.
+RerollPlan extendCandidate(ArrayRef ops,
+ const llvm::DenseMap &indexOf, RerollPlan plan)
+{
+ while (plan.cand.start >= plan.cand.period) {
+ Candidate c = plan.cand;
+ c.start -= c.period;
+ c.count += 1;
+ auto extended = verifyCandidate(ops, indexOf, c);
+ if (!extended) {
+ break;
+ }
+ plan = *extended;
+ }
+ while (true) {
+ Candidate c = plan.cand;
+ c.count += 1;
+ auto extended = verifyCandidate(ops, indexOf, c);
+ if (!extended) {
+ break;
+ }
+ plan = *extended;
+ }
+ return plan;
+}
+
+//===----------------------------------------------------------------------===//
+// Materialization
+//===----------------------------------------------------------------------===//
+
+void materialize(ArrayRef ops, const RerollPlan &plan)
+{
+ size_t start = plan.cand.start, p = plan.cand.period, count = plan.cand.count;
+ Operation *first = ops[start];
+ Location loc = first->getLoc();
+ OpBuilder builder(first);
+
+ Value lb = arith::ConstantIndexOp::create(builder, loc, 0);
+ Value ub = arith::ConstantIndexOp::create(builder, loc, count);
+ Value step = arith::ConstantIndexOp::create(builder, loc, 1);
+
+ auto forOp = scf::ForOp::create(builder, loc, lb, ub, step, plan.inits);
+ Block *body = forOp.getBody();
+ builder.setInsertionPointToStart(body);
+
+ // Clone window 0 as the loop body, remapping operands per classification.
+ SmallVector cloned(p);
+ IRMapping mapping; // within-window result mapping
+ for (size_t j = 0; j < p; ++j) {
+ Operation *proto = ops[start + j];
+ Operation *clone = proto->cloneWithoutRegions(mapping);
+ // Fix up operands that are threaded from the previous iteration: the
+ // prototype (window 0) uses the init values there.
+ for (unsigned t = 0; t < clone->getNumOperands(); ++t) {
+ auto it = plan.threadedUse.find({(unsigned)j, t});
+ if (it != plan.threadedUse.end()) {
+ clone->setOperand(t, forOp.getRegionIterArg(it->second));
+ }
+ }
+ builder.insert(clone);
+ cloned[j] = clone;
+ }
+ SmallVector yields;
+ for (auto [defPos, resNo] : plan.slots) {
+ yields.push_back(cloned[defPos]->getResult(resNo));
+ }
+ scf::YieldOp::create(builder, loc, yields);
+
+ // Rewire external uses of the last window's results to the loop results.
+ for (const auto &[slotIdx, slot] : llvm::enumerate(plan.slots)) {
+ auto [defPos, resNo] = slot;
+ Operation *lastOp = ops[start + (count - 1) * p + defPos];
+ lastOp->getResult(resNo).replaceAllUsesWith(forOp.getResult(slotIdx));
+ }
+
+ // Erase the original ops, last first (uses before defs).
+ for (size_t i = start + count * p; i-- > start;) {
+ ops[i]->dropAllUses();
+ ops[i]->erase();
+ }
+}
+
+//===----------------------------------------------------------------------===//
+// Driver
+//===----------------------------------------------------------------------===//
+
+bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
+{
+ SmallVector ops;
+ for (Operation &op : block.without_terminator()) {
+ ops.push_back(&op);
+ }
+ if (ops.size() < 2 * minPeriod) {
+ return false;
+ }
+
+ llvm::DenseMap indexOf;
+ for (const auto &[i, op] : llvm::enumerate(ops)) {
+ indexOf[op] = i;
+ }
+
+ SmallVector hashes;
+ computeHashes(ops, indexOf, hashes);
+
+ SmallVector candidates;
+ findCandidates(hashes, minPeriod, /*maxPeriods=*/16, candidates);
+
+ // Verify, extend, and pick non-overlapping plans greedily by savings.
+ SmallVector plans;
+ for (Candidate cand : candidates) {
+ auto plan = verifyCandidate(ops, indexOf, cand);
+ if (!plan) {
+ continue;
+ }
+ *plan = extendCandidate(ops, indexOf, *plan);
+ if ((plan->cand.count - 1) * plan->cand.period >= minSavings) {
+ plans.push_back(std::move(*plan));
+ }
+ }
+ llvm::sort(plans, [](const RerollPlan &a, const RerollPlan &b) {
+ return (a.cand.count - 1) * a.cand.period > (b.cand.count - 1) * b.cand.period;
+ });
+
+ SmallVector> used;
+ SmallVector accepted;
+ for (const RerollPlan &plan : plans) {
+ size_t s = plan.cand.start, e = s + plan.cand.count * plan.cand.period;
+ bool overlaps = llvm::any_of(
+ used, [&](auto range) { return s < range.second && range.first < e; });
+ if (!overlaps) {
+ used.push_back({s, e});
+ accepted.push_back(&plan);
+ }
+ }
+
+ // Materialize from the highest start index down so op indices of pending
+ // plans remain valid.
+ llvm::sort(accepted, [](const RerollPlan *a, const RerollPlan *b) {
+ return a->cand.start > b->cand.start;
+ });
+ for (const RerollPlan *plan : accepted) {
+ LLVM_DEBUG(dbgs() << "rerolling: start=" << plan->cand.start
+ << " period=" << plan->cand.period
+ << " count=" << plan->cand.count
+ << " slots=" << plan->slots.size() << "\n");
+ materialize(ops, *plan);
+ }
+ return !accepted.empty();
+}
+
+} // namespace
+
+namespace catalyst {
+
+#define GEN_PASS_DECL_REROLLLOOPSPASS
+#define GEN_PASS_DEF_REROLLLOOPSPASS
+#include "Catalyst/Transforms/Passes.h.inc"
+
+struct RerollLoopsPass : public impl::RerollLoopsPassBase {
+ using impl::RerollLoopsPassBase::RerollLoopsPassBase;
+
+ void runOnOperation() override
+ {
+ // Iterate to a fixpoint (bounded): rerolling creates new blocks (loop
+ // bodies) that may contain further repeats, e.g. nested loops.
+ bool changed = true;
+ unsigned rounds = 0;
+ while (changed && rounds++ < 4) {
+ changed = false;
+ SmallVector blocks;
+ getOperation()->walk([&](Block *block) { blocks.push_back(block); });
+ for (Block *block : blocks) {
+ changed |= processBlock(*block, minPeriod, minSavings);
+ }
+ }
+ }
+};
+
+} // namespace catalyst
diff --git a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
new file mode 100644
index 0000000000..ef0c7665e9
--- /dev/null
+++ b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
@@ -0,0 +1,283 @@
+// Copyright 2026 Xanadu Quantum Technologies Inc.
+
+// Licensed under the Apache License, Version 2.0 (the "License");
+// you may not use this file except in compliance with the License.
+// You may obtain a copy of the License at
+
+// http://www.apache.org/licenses/LICENSE-2.0
+
+// Unless required by applicable law or agreed to in writing, software
+// distributed under the License is distributed on an "AS IS" BASIS,
+// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+// See the License for the specific language governing permissions and
+// limitations under the License.
+
+#define DEBUG_TYPE "scalarize-tensor-extracts"
+
+#include "mlir/Dialect/Arith/IR/Arith.h"
+#include "mlir/Dialect/Linalg/IR/Linalg.h"
+#include "mlir/Dialect/Tensor/IR/Tensor.h"
+#include "mlir/IR/IRMapping.h"
+#include "mlir/Interfaces/SideEffectInterfaces.h"
+#include "mlir/Pass/Pass.h"
+#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
+
+using namespace llvm;
+using namespace mlir;
+
+// Traced programs that compute quantum gate parameters from runtime inputs (e.g.
+// Trotterized Hamiltonians with runtime coefficients, see catalyst issue #2759)
+// produce long chains of *small* tensor ops whose only consumers are scalar
+// `tensor.extract` operations feeding gate parameters. Each such tensor survives
+// bufferization as an allocation plus copies, and each surviving `linalg.generic`
+// is later unrolled into loops, amplifying the IR by orders of magnitude.
+//
+// The patterns in this pass sink `tensor.extract` through the producers of such
+// tensors, computing the extracted element directly in scalar arithmetic:
+//
+// * extract(linalg.generic) -> inline the generic's scalar payload
+// * extract(tensor.collapse_shape) -> extract from the source with expanded indices
+// * extract(tensor.extract_slice) -> extract from the source with offset indices
+//
+// Applied to a fixpoint, extraction sinks to the leaves of the dataflow, the
+// intermediate tensors become dead, and the scalar chains remain (deduplicated by
+// a follow-up CSE). To avoid code growth on large tensors, generic-payload
+// inlining is restricted to small statically-shaped results.
+
+namespace {
+
+/// Upper bound on the number of elements of a tensor whose producer payload we
+/// are willing to clone per extraction site.
+constexpr int64_t kMaxScalarizedElements = 16;
+
+/// Fold tensor.extract(linalg.generic) by inlining the generic's scalar payload
+/// at the extraction point, for elementwise (all-parallel) generics.
+struct ExtractOfGeneric : public OpRewritePattern {
+ using OpRewritePattern::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
+ PatternRewriter &rewriter) const override
+ {
+ auto genericOp = extractOp.getTensor().getDefiningOp();
+ if (!genericOp) {
+ return failure();
+ }
+
+ // Only elementwise generics: every iterator is parallel.
+ if (genericOp.getNumParallelLoops() != genericOp.getNumLoops()) {
+ return failure();
+ }
+
+ // Restrict to small statically shaped results to bound code growth.
+ auto resultType = dyn_cast(extractOp.getTensor().getType());
+ if (!resultType || !resultType.hasStaticShape() ||
+ resultType.getNumElements() > kMaxScalarizedElements) {
+ return failure();
+ }
+
+ // Identify which result of the generic is being extracted and require an
+ // identity indexing map for it, so the iteration indices equal the
+ // extraction indices.
+ auto resultNumber = cast(extractOp.getTensor()).getResultNumber();
+ OpOperand *initOperand = genericOp.getDpsInitOperand(resultNumber);
+ AffineMap outputMap = genericOp.getMatchingIndexingMap(initOperand);
+ if (!outputMap.isIdentity()) {
+ return failure();
+ }
+
+ Block *body = genericOp.getBody();
+
+ // The payload must be speculatable scalar code and must not read the
+ // accumulator (output block argument).
+ for (Operation &op : body->without_terminator()) {
+ if (!isPure(&op)) {
+ return failure();
+ }
+ }
+ for (OpOperand &outOperand : genericOp.getDpsInitsMutable()) {
+ BlockArgument outArg = body->getArgument(outOperand.getOperandNumber());
+ if (!outArg.use_empty()) {
+ return failure();
+ }
+ }
+
+ Location loc = extractOp.getLoc();
+ SmallVector iterIndices(extractOp.getIndices());
+
+ // Materialize scalar operands: one tensor.extract per generic input, at
+ // indices given by composing that input's indexing map with the
+ // extraction indices.
+ IRMapping mapping;
+ for (OpOperand *inOperand : genericOp.getDpsInputOperands()) {
+ BlockArgument blockArg = body->getArgument(inOperand->getOperandNumber());
+ Value input = inOperand->get();
+ if (!isa(input.getType())) {
+ // Scalar operands of the generic map through unchanged.
+ mapping.map(blockArg, input);
+ continue;
+ }
+ AffineMap inputMap = genericOp.getMatchingIndexingMap(inOperand);
+ SmallVector inputIndices;
+ for (AffineExpr expr : inputMap.getResults()) {
+ if (auto dimExpr = dyn_cast(expr)) {
+ inputIndices.push_back(iterIndices[dimExpr.getPosition()]);
+ }
+ else if (auto constExpr = dyn_cast(expr)) {
+ inputIndices.push_back(arith::ConstantIndexOp::create(
+ rewriter, loc, constExpr.getValue()));
+ }
+ else {
+ return failure();
+ }
+ }
+ Value scalar = tensor::ExtractOp::create(rewriter, loc, input, inputIndices);
+ mapping.map(blockArg, scalar);
+ }
+
+ // Clone the payload, resolving linalg.index to the extraction indices.
+ for (Operation &op : body->without_terminator()) {
+ if (auto indexOp = dyn_cast(op)) {
+ mapping.map(indexOp.getResult(), iterIndices[indexOp.getDim()]);
+ continue;
+ }
+ rewriter.clone(op, mapping);
+ }
+
+ auto yieldOp = cast(body->getTerminator());
+ Value result = mapping.lookupOrDefault(yieldOp.getOperand(resultNumber));
+ rewriter.replaceOp(extractOp, result);
+ return success();
+ }
+};
+
+/// Fold tensor.extract(tensor.collapse_shape) for collapses that only drop or
+/// merge unit dimensions (at most one non-unit dimension per reassociation
+/// group), by extracting directly from the source.
+struct ExtractOfCollapseShape : public OpRewritePattern {
+ using OpRewritePattern::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
+ PatternRewriter &rewriter) const override
+ {
+ auto collapseOp = extractOp.getTensor().getDefiningOp();
+ if (!collapseOp) {
+ return failure();
+ }
+ auto srcType = collapseOp.getSrcType();
+ if (!srcType.hasStaticShape()) {
+ return failure();
+ }
+
+ Location loc = extractOp.getLoc();
+ SmallVector srcIndices(srcType.getRank());
+
+ Value zero;
+ auto getZero = [&]() {
+ if (!zero) {
+ zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
+ }
+ return zero;
+ };
+
+ SmallVector groups = collapseOp.getReassociationIndices();
+ // A rank-0 result means every source dimension is a unit dimension.
+ if (groups.empty()) {
+ for (int64_t dim = 0; dim < srcType.getRank(); ++dim) {
+ srcIndices[dim] = getZero();
+ }
+ }
+ for (const auto &[groupIdx, group] : llvm::enumerate(groups)) {
+ int64_t nonUnitDim = -1;
+ for (int64_t srcDim : group) {
+ if (srcType.getDimSize(srcDim) != 1) {
+ if (nonUnitDim != -1) {
+ return failure(); // true merge of two non-unit dims
+ }
+ nonUnitDim = srcDim;
+ }
+ }
+ for (int64_t srcDim : group) {
+ srcIndices[srcDim] = (srcDim == nonUnitDim)
+ ? extractOp.getIndices()[groupIdx]
+ : getZero();
+ }
+ }
+
+ rewriter.replaceOpWithNewOp(extractOp, collapseOp.getSrc(),
+ srcIndices);
+ return success();
+ }
+};
+
+/// Fold tensor.extract(tensor.extract_slice) by extracting from the source at
+/// offset + index * stride.
+struct ExtractOfExtractSlice : public OpRewritePattern {
+ using OpRewritePattern::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
+ PatternRewriter &rewriter) const override
+ {
+ auto sliceOp = extractOp.getTensor().getDefiningOp();
+ if (!sliceOp) {
+ return failure();
+ }
+
+ Location loc = extractOp.getLoc();
+ int64_t srcRank = sliceOp.getSourceType().getRank();
+
+ // The slice may be rank-reducing: map each source dim to its position in
+ // the result (or none if the dim was dropped).
+ llvm::SmallBitVector droppedDims = sliceOp.getDroppedDims();
+
+ auto materialize = [&](OpFoldResult ofr) -> Value {
+ if (auto val = dyn_cast(ofr)) {
+ return val;
+ }
+ return arith::ConstantIndexOp::create(
+ rewriter, loc, cast(cast(ofr)).getInt());
+ };
+
+ SmallVector srcIndices;
+ unsigned resultDim = 0;
+ for (int64_t dim = 0; dim < srcRank; ++dim) {
+ Value offset = materialize(sliceOp.getMixedOffsets()[dim]);
+ if (droppedDims.test(dim)) {
+ srcIndices.push_back(offset);
+ continue;
+ }
+ Value index = extractOp.getIndices()[resultDim++];
+ Value stride = materialize(sliceOp.getMixedStrides()[dim]);
+ Value scaled = arith::MulIOp::create(rewriter, loc, index, stride);
+ srcIndices.push_back(arith::AddIOp::create(rewriter, loc, offset, scaled));
+ }
+
+ rewriter.replaceOpWithNewOp(extractOp, sliceOp.getSource(),
+ srcIndices);
+ return success();
+ }
+};
+
+} // namespace
+
+namespace catalyst {
+
+#define GEN_PASS_DEF_SCALARIZETENSOREXTRACTSPASS
+#include "Catalyst/Transforms/Passes.h.inc"
+
+struct ScalarizeTensorExtractsPass
+ : public impl::ScalarizeTensorExtractsPassBase {
+ using impl::ScalarizeTensorExtractsPassBase<
+ ScalarizeTensorExtractsPass>::ScalarizeTensorExtractsPassBase;
+
+ void runOnOperation() override
+ {
+ MLIRContext *context = &getContext();
+ RewritePatternSet patterns(context);
+ patterns.add(context);
+ if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) {
+ signalPassFailure();
+ }
+ }
+};
+
+} // namespace catalyst
diff --git a/mlir/test/Catalyst/RerollLoopsTest.mlir b/mlir/test/Catalyst/RerollLoopsTest.mlir
new file mode 100644
index 0000000000..28a59dd007
--- /dev/null
+++ b/mlir/test/Catalyst/RerollLoopsTest.mlir
@@ -0,0 +1,133 @@
+// Copyright 2026 Xanadu Quantum Technologies Inc.
+
+// Licensed under the Apache License, Version 2.0 (the "License");
+// you may not use this file except in compliance with the License.
+// You may obtain a copy of the License at
+
+// http://www.apache.org/licenses/LICENSE-2.0
+
+// Unless required by applicable law or agreed to in writing, software
+// distributed under the License is distributed on an "AS IS" BASIS,
+// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+// See the License for the specific language governing permissions and
+// limitations under the License.
+
+// RUN: quantum-opt --reroll-loops="min-period=2 min-savings=4" --split-input-file %s | FileCheck %s
+
+// A scalar chain of alternating ops rerolls into an scf.for threading one
+// value. The first addf seeds the loop (its operand pattern differs), the
+// trailing mulf stays as an epilogue.
+
+// CHECK-LABEL: @scalar_chain
+// CHECK: %[[SEED:.+]] = arith.addf
+// CHECK: %[[FOR:.+]] = scf.for {{.*}} iter_args(%[[IT:.+]] = %[[SEED]]) -> (f64)
+// CHECK: %[[M:.+]] = arith.mulf %[[IT]],
+// CHECK: %[[A:.+]] = arith.addf %[[M]],
+// CHECK: scf.yield %[[A]] : f64
+// CHECK: %[[EPI:.+]] = arith.mulf %[[FOR]],
+// CHECK: return %[[EPI]]
+func.func @scalar_chain(%arg0: f64, %c: f64) -> f64 {
+ %0 = arith.addf %arg0, %c : f64
+ %1 = arith.mulf %0, %c : f64
+ %2 = arith.addf %1, %c : f64
+ %3 = arith.mulf %2, %c : f64
+ %4 = arith.addf %3, %c : f64
+ %5 = arith.mulf %4, %c : f64
+ %6 = arith.addf %5, %c : f64
+ %7 = arith.mulf %6, %c : f64
+ %8 = arith.addf %7, %c : f64
+ %9 = arith.mulf %8, %c : f64
+ %10 = arith.addf %9, %c : f64
+ %11 = arith.mulf %10, %c : f64
+ %12 = arith.addf %11, %c : f64
+ %13 = arith.mulf %12, %c : f64
+ %14 = arith.addf %13, %c : f64
+ %15 = arith.mulf %14, %c : f64
+ return %15 : f64
+}
+
+// -----
+
+// A repeated gate sequence threading two qubits rerolls with both qubit
+// values as iter_args; the rotation angle is loop-invariant. The pass may pick
+// any rotation of the repeated window (here the run starts at the first CNOT),
+// leaving a prologue/epilogue outside the loop.
+
+// CHECK-LABEL: @gate_sequence
+// CHECK: quantum.alloc
+// CHECK: %[[FOR:.+]]:2 = scf.for {{.*}} iter_args(%[[Q0:.+]] = %{{.+}}, %[[Q1:.+]] = %{{.+}}) -> (!quantum.bit, !quantum.bit)
+// CHECK: %[[CNOT:.+]]:2 = quantum.custom "CNOT"() %[[Q0]], %[[Q1]]
+// CHECK: %[[H:.+]] = quantum.custom "Hadamard"() %[[CNOT]]#0
+// CHECK: %[[RZ:.+]] = quantum.custom "RZ"(%{{.+}}) %[[CNOT]]#1
+// CHECK: scf.yield %[[H]], %[[RZ]]
+// CHECK: %[[LAST:.+]]:2 = quantum.custom "CNOT"() %[[FOR]]#0, %[[FOR]]#1
+// CHECK: quantum.insert %{{.+}}[ 0], %[[LAST]]#0
+func.func @gate_sequence(%theta: f64) -> !quantum.reg {
+ %r0 = quantum.alloc( 2) : !quantum.reg
+ %q0 = quantum.extract %r0[ 0] : !quantum.reg -> !quantum.bit
+ %q1 = quantum.extract %r0[ 1] : !quantum.reg -> !quantum.bit
+
+ %h0 = quantum.custom "Hadamard"() %q0 : !quantum.bit
+ %z0 = quantum.custom "RZ"(%theta) %q1 : !quantum.bit
+ %c0:2 = quantum.custom "CNOT"() %h0, %z0 : !quantum.bit, !quantum.bit
+
+ %h1 = quantum.custom "Hadamard"() %c0#0 : !quantum.bit
+ %z1 = quantum.custom "RZ"(%theta) %c0#1 : !quantum.bit
+ %c1:2 = quantum.custom "CNOT"() %h1, %z1 : !quantum.bit, !quantum.bit
+
+ %h2 = quantum.custom "Hadamard"() %c1#0 : !quantum.bit
+ %z2 = quantum.custom "RZ"(%theta) %c1#1 : !quantum.bit
+ %c2:2 = quantum.custom "CNOT"() %h2, %z2 : !quantum.bit, !quantum.bit
+
+ %h3 = quantum.custom "Hadamard"() %c2#0 : !quantum.bit
+ %z3 = quantum.custom "RZ"(%theta) %c2#1 : !quantum.bit
+ %c3:2 = quantum.custom "CNOT"() %h3, %z3 : !quantum.bit, !quantum.bit
+
+ %r1 = quantum.insert %r0[ 0], %c3#0 : !quantum.reg, !quantum.bit
+ %r2 = quantum.insert %r1[ 1], %c3#1 : !quantum.reg, !quantum.bit
+ return %r2 : !quantum.reg
+}
+
+// -----
+
+// Iterations with *different* angle values must not be rerolled: the varying
+// operand fails loop-invariance verification.
+
+// CHECK-LABEL: @varying_angles
+// CHECK-NOT: scf.for
+func.func @varying_angles(%t0: f64, %t1: f64, %t2: f64, %t3: f64) -> !quantum.reg {
+ %r0 = quantum.alloc( 1) : !quantum.reg
+ %q0 = quantum.extract %r0[ 0] : !quantum.reg -> !quantum.bit
+ %a = quantum.custom "RZ"(%t0) %q0 : !quantum.bit
+ %b = quantum.custom "RX"(%t0) %a : !quantum.bit
+ %c = quantum.custom "RZ"(%t1) %b : !quantum.bit
+ %d = quantum.custom "RX"(%t1) %c : !quantum.bit
+ %e = quantum.custom "RZ"(%t2) %d : !quantum.bit
+ %f = quantum.custom "RX"(%t2) %e : !quantum.bit
+ %g = quantum.custom "RZ"(%t3) %f : !quantum.bit
+ %h = quantum.custom "RX"(%t3) %g : !quantum.bit
+ %r1 = quantum.insert %r0[ 0], %h : !quantum.reg, !quantum.bit
+ return %r1 : !quantum.reg
+}
+
+// -----
+
+// Results used outside the repeat (other than the final threaded values) block
+// rerolling: here every intermediate feeds the final sum.
+
+// CHECK-LABEL: @external_uses
+// CHECK-NOT: scf.for
+func.func @external_uses(%arg0: f64, %c: f64) -> f64 {
+ %0 = arith.addf %arg0, %c : f64
+ %1 = arith.mulf %0, %c : f64
+ %2 = arith.addf %1, %c : f64
+ %3 = arith.mulf %2, %c : f64
+ %4 = arith.addf %3, %c : f64
+ %5 = arith.mulf %4, %c : f64
+ %6 = arith.addf %5, %c : f64
+ %7 = arith.mulf %6, %c : f64
+ %s0 = arith.addf %1, %3 : f64
+ %s1 = arith.addf %s0, %5 : f64
+ %s2 = arith.addf %s1, %7 : f64
+ return %s2 : f64
+}
diff --git a/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir b/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
new file mode 100644
index 0000000000..d41df225da
--- /dev/null
+++ b/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
@@ -0,0 +1,146 @@
+// Copyright 2026 Xanadu Quantum Technologies Inc.
+
+// Licensed under the Apache License, Version 2.0 (the "License");
+// you may not use this file except in compliance with the License.
+// You may obtain a copy of the License at
+
+// http://www.apache.org/licenses/LICENSE-2.0
+
+// Unless required by applicable law or agreed to in writing, software
+// distributed under the License is distributed on an "AS IS" BASIS,
+// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+// See the License for the specific language governing permissions and
+// limitations under the License.
+
+// RUN: quantum-opt --scalarize-tensor-extracts --canonicalize --split-input-file %s | FileCheck %s
+
+// The extract_slice + collapse_shape + extract chain produced when tracing
+// indexes a 1-D runtime tensor collapses to a single extract.
+
+// CHECK-LABEL: @slice_collapse_extract
+// CHECK-SAME: (%[[ARG:.+]]: tensor<15xf64>)
+// CHECK: %[[C14:.+]] = arith.constant 14 : index
+// CHECK: %[[RES:.+]] = tensor.extract %[[ARG]][%[[C14]]] : tensor<15xf64>
+// CHECK-NOT: tensor.extract_slice
+// CHECK-NOT: tensor.collapse_shape
+// CHECK: return %[[RES]]
+func.func @slice_collapse_extract(%arg0: tensor<15xf64>) -> f64 {
+ %s = tensor.extract_slice %arg0[14] [1] [1] : tensor<15xf64> to tensor<1xf64>
+ %c = tensor.collapse_shape %s [] : tensor<1xf64> into tensor
+ %e = tensor.extract %c[] : tensor
+ return %e : f64
+}
+
+// -----
+
+// Extracting one element of an elementwise linalg.generic inlines the scalar
+// payload; the generic and its tensor.empty become dead.
+
+// CHECK-LABEL: @extract_of_generic
+// CHECK-SAME: (%[[A:.+]]: tensor<2x2xf64>, %[[B:.+]]: tensor<2x2xf64>)
+// CHECK-DAG: %[[C0:.+]] = arith.constant 0 : index
+// CHECK-DAG: %[[C1:.+]] = arith.constant 1 : index
+// CHECK-DAG: %[[EA:.+]] = tensor.extract %[[A]][%[[C0]], %[[C1]]]
+// CHECK-DAG: %[[EB:.+]] = tensor.extract %[[B]][%[[C0]], %[[C1]]]
+// CHECK: %[[RES:.+]] = arith.mulf %[[EA]], %[[EB]] : f64
+// CHECK-NOT: linalg.generic
+// CHECK: return %[[RES]]
+#map = affine_map<(d0, d1) -> (d0, d1)>
+func.func @extract_of_generic(%arg0: tensor<2x2xf64>, %arg1: tensor<2x2xf64>) -> f64 {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %empty = tensor.empty() : tensor<2x2xf64>
+ %prod = linalg.generic {indexing_maps = [#map, #map, #map], iterator_types = ["parallel", "parallel"]}
+ ins(%arg0, %arg1 : tensor<2x2xf64>, tensor<2x2xf64>) outs(%empty : tensor<2x2xf64>) {
+ ^bb0(%in0: f64, %in1: f64, %out: f64):
+ %m = arith.mulf %in0, %in1 : f64
+ linalg.yield %m : f64
+ } -> tensor<2x2xf64>
+ %res = tensor.extract %prod[%c0, %c1] : tensor<2x2xf64>
+ return %res : f64
+}
+
+// -----
+
+// Broadcast (rank-0 to 2x2) generics fold to an extract of the rank-0 source.
+
+// CHECK-LABEL: @extract_of_broadcast
+// CHECK-SAME: (%[[A:.+]]: tensor)
+// CHECK: %[[E:.+]] = tensor.extract %[[A]][] : tensor
+// CHECK-NOT: linalg.generic
+// CHECK: return %[[E]]
+#map0 = affine_map<(d0, d1) -> ()>
+#map1 = affine_map<(d0, d1) -> (d0, d1)>
+func.func @extract_of_broadcast(%arg0: tensor) -> f64 {
+ %c1 = arith.constant 1 : index
+ %empty = tensor.empty() : tensor<2x2xf64>
+ %bcast = linalg.generic {indexing_maps = [#map0, #map1], iterator_types = ["parallel", "parallel"]}
+ ins(%arg0 : tensor) outs(%empty : tensor<2x2xf64>) {
+ ^bb0(%in: f64, %out: f64):
+ linalg.yield %in : f64
+ } -> tensor<2x2xf64>
+ %res = tensor.extract %bcast[%c1, %c1] : tensor<2x2xf64>
+ return %res : f64
+}
+
+// -----
+
+// Reductions must not be scalarized: the payload reads the accumulator.
+
+// CHECK-LABEL: @reduction_untouched
+// CHECK: linalg.generic
+// CHECK: tensor.extract
+#map_in = affine_map<(d0) -> (d0)>
+#map_out = affine_map<(d0) -> ()>
+func.func @reduction_untouched(%arg0: tensor<8xf64>) -> f64 {
+ %cst = arith.constant 0.0 : f64
+ %empty = tensor.empty() : tensor
+ %fill = linalg.fill ins(%cst : f64) outs(%empty : tensor) -> tensor
+ %sum = linalg.generic {indexing_maps = [#map_in, #map_out], iterator_types = ["reduction"]}
+ ins(%arg0 : tensor<8xf64>) outs(%fill : tensor) {
+ ^bb0(%in: f64, %acc: f64):
+ %a = arith.addf %in, %acc : f64
+ linalg.yield %a : f64
+ } -> tensor
+ %res = tensor.extract %sum[] : tensor
+ return %res : f64
+}
+
+// -----
+
+// Large tensors must not be scalarized (payload cloning is capped).
+
+// CHECK-LABEL: @large_tensor_untouched
+// CHECK: linalg.generic
+// CHECK: tensor.extract
+#map2 = affine_map<(d0) -> (d0)>
+func.func @large_tensor_untouched(%arg0: tensor<100xf64>) -> f64 {
+ %c5 = arith.constant 5 : index
+ %empty = tensor.empty() : tensor<100xf64>
+ %sq = linalg.generic {indexing_maps = [#map2, #map2], iterator_types = ["parallel"]}
+ ins(%arg0 : tensor<100xf64>) outs(%empty : tensor<100xf64>) {
+ ^bb0(%in: f64, %out: f64):
+ %m = arith.mulf %in, %in : f64
+ linalg.yield %m : f64
+ } -> tensor<100xf64>
+ %res = tensor.extract %sq[%c5] : tensor<100xf64>
+ return %res : f64
+}
+
+// -----
+
+// Rank-reducing extract_slice: index arithmetic offset + i * stride.
+
+// CHECK-LABEL: @strided_slice_extract
+// CHECK-SAME: (%[[ARG:.+]]: tensor<4x6xf64>, %[[I:.+]]: index)
+// CHECK-DAG: %[[C2:.+]] = arith.constant 2 : index
+// CHECK-DAG: %[[C1:.+]] = arith.constant 1 : index
+// CHECK-DAG: %[[SCALED:.+]] = arith.muli %[[I]], %[[C2]]
+// CHECK-DAG: %[[COL:.+]] = arith.addi %[[SCALED]], %[[C1]]
+// CHECK: %[[RES:.+]] = tensor.extract %[[ARG]][%[[C2]], %[[COL]]] : tensor<4x6xf64>
+// CHECK: return %[[RES]]
+func.func @strided_slice_extract(%arg0: tensor<4x6xf64>, %i: index) -> f64 {
+ %s = tensor.extract_slice %arg0[2, 1] [1, 3] [1, 2] : tensor<4x6xf64> to tensor<3xf64>
+ %e = tensor.extract %s[%i] : tensor<3xf64>
+ return %e : f64
+}
From 62b057c94be9f951b9f7af58b29f177581417278 Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Sun, 12 Jul 2026 02:19:14 -0400
Subject: [PATCH 2/9] Fix BufferizationStage IR amplification for
runtime-coefficient Hamiltonians
Add scalarize-tensor-extracts and reroll-loops passes and re-order the
default pipeline so runtime-coefficient Trotterization no longer amplifies
IR through bufferization (9.9x -> ~1.0x on the H2 QPE benchmark).
Fixes #2759
---
doc/releases/changelog-dev.md | 34 ++++-----
.../pytest/test_trotter_runtime_coeffs.py | 6 +-
mlir/include/Catalyst/Transforms/Passes.td | 5 +-
.../DefaultPipelines/DefaultPipelines.h | 10 +--
.../Catalyst/Transforms/RerollLoopsPass.cpp | 73 ++++---------------
.../ScalarizeTensorExtractsPass.cpp | 27 +++----
6 files changed, 50 insertions(+), 105 deletions(-)
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index c0b532c3c4..ead45f46d7 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -12,23 +12,6 @@
Improvements ðŸ›
-* IR amplification for circuits whose gate parameters are computed from runtime values
- (e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
- Three changes work together in the default pipeline:
- the new `scalarize-tensor-extracts` pass sinks scalar `tensor.extract` operations through
- `tensor.extract_slice`, `tensor.collapse_shape`, and small elementwise `linalg.generic`
- producers so gate-angle dataflow becomes pure scalar arithmetic instead of thousands of
- tiny tensors that each survive bufferization as an allocation and copies;
- `linalg-fuse-elementwise-ops` plus additional `cse` applications (before and after
- `one-shot-bufferize`) remove the duplicate angle computations produced by tracing; and
- the new `reroll-loops` pass reconstructs the loops that Python tracing unrolled, by
- detecting tandem repeats of structurally isomorphic operation windows (via structural
- hashing), verifying that cross-window dataflow is limited to threaded SSA values (such as
- qubit values) plus loop-invariant values, and replacing each repeat with an `scf.for`.
- For the H2 QPE benchmark from the issue, gate-op volume after HLO lowering drops about
- 5x and downstream IR, compile time, and peak memory drop accordingly.
- [(#2759)](https://github.com/PennyLaneAI/catalyst/issues/2759)
-
* The `decompose-lowering` pass now supports applying a selection of the available decomposition rules via the `target_rules` parameter.
The pass also no longer applies the `inline`, `cse` and `canonicalize` passes to avoid unnecessary IR mutations.
Instead, decomposition rules are deterministically inlined by a custom function (`inline` is non-deterministic, using an estimated benefit and threshold as criteria for inlining).
@@ -176,6 +159,23 @@
PennyLane.
[(#2769)](https://github.com/PennyLaneAI/catalyst/pull/2769)
+* IR amplification for circuits whose gate parameters are computed from runtime values
+ (e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
+ Three changes work together in the default pipeline:
+ the new `scalarize-tensor-extracts` pass sinks scalar `tensor.extract` operations through
+ `tensor.extract_slice`, `tensor.collapse_shape`, and small elementwise `linalg.generic`
+ producers so gate-angle dataflow becomes pure scalar arithmetic instead of thousands of
+ tiny tensors that each survive bufferization as an allocation and copies;
+ `linalg-fuse-elementwise-ops` plus additional `cse` applications (before and after
+ `one-shot-bufferize`) remove the duplicate angle computations produced by tracing; and
+ the new `reroll-loops` pass reconstructs the loops that Python tracing unrolled, by
+ detecting tandem repeats of structurally isomorphic operation windows (via structural
+ hashing), verifying that cross-window dataflow is limited to threaded SSA values (such as
+ qubit values) plus loop-invariant values, and replacing each repeat with an `scf.for`.
+ For a representative Trotterized QPE workload with runtime coefficients, gate-op volume
+ after HLO lowering drops about 5x and downstream IR, compile time, and peak memory drop
+ accordingly.
+
Breaking changes 💔
* Catalyst's xDSL dependencies have been updated to `xdsl` 0.63.0 and `xdsl-jax` 0.5.2.
diff --git a/frontend/test/pytest/test_trotter_runtime_coeffs.py b/frontend/test/pytest/test_trotter_runtime_coeffs.py
index 15f446a5dd..b71720589b 100644
--- a/frontend/test/pytest/test_trotter_runtime_coeffs.py
+++ b/frontend/test/pytest/test_trotter_runtime_coeffs.py
@@ -13,9 +13,9 @@
# limitations under the License.
"""Integration tests for the IR-amplification fixes for runtime-coefficient
-Hamiltonians (issue #2759): scalarize-tensor-extracts, elementwise fusion, and
-reroll-loops in the default pipeline must preserve numerics for Trotterized
-workloads with runtime coefficients."""
+Hamiltonians: scalarize-tensor-extracts, elementwise fusion, and reroll-loops
+in the default pipeline must preserve numerics for Trotterized workloads with
+runtime coefficients."""
import numpy as np
import pennylane as qml
diff --git a/mlir/include/Catalyst/Transforms/Passes.td b/mlir/include/Catalyst/Transforms/Passes.td
index 30eddd487f..37ab0667fd 100644
--- a/mlir/include/Catalyst/Transforms/Passes.td
+++ b/mlir/include/Catalyst/Transforms/Passes.td
@@ -338,8 +338,7 @@ def RerollLoopsPass : Pass<"reroll-loops"> {
let description = [{
Tracing a Python program unrolls its loops: N structurally identical
circuit segments (Trotter steps, layers, folds) arrive as N copies of the
- same op sequence, and every downstream stage amplifies each copy (see
- issue #2759).
+ same op sequence, and every downstream stage amplifies each copy.
This pass detects maximal tandem repeats of structurally isomorphic op
windows via shift-invariant structural hashing, verifies that
@@ -381,7 +380,7 @@ def ScalarizeTensorExtractsPass : Pass<"scalarize-tensor-extracts"> {
small tensor operations whose only consumers are scalar `tensor.extract`
operations. Each intermediate tensor survives bufferization as an allocation
with copies and each `linalg.generic` is later unrolled into loops, causing
- severe IR amplification (see issue #2759).
+ severe IR amplification.
This pass folds `tensor.extract` through elementwise `linalg.generic` (by
inlining the scalar payload), `tensor.collapse_shape`, and
diff --git a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
index d2d91880a9..7a308085f1 100644
--- a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
+++ b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
@@ -73,22 +73,22 @@ const PipelineList pipelineList{
"cse",
// Sink scalar extractions through small-tensor producers and fuse the
// remaining elementwise ops. Traced gate-parameter dataflow (e.g. runtime
- // Hamiltonian coefficients, issue #2759) otherwise reaches bufferization
- // as thousands of tiny tensor ops that each become an alloc + copy.
+ // Hamiltonian coefficients) otherwise reaches bufferization as thousands
+ // of tiny tensor ops that each become an alloc + copy.
"scalarize-tensor-extracts",
"func.func(linalg-fuse-elementwise-ops)",
"canonicalize",
"cse",
// Reconstruct the loops that tracing unrolled (e.g. Trotter steps); a
// repeat of multiplicity k shrinks its region k-fold before the
- // bufferization and LLVM stages amplify it (issue #2759).
+ // bufferization and LLVM stages amplify it.
"reroll-loops",
"func.func(linalg-detensorize{aggressive-mode})",
"detensorize-scf",
"detensorize-function-boundary",
// Detensorization is what materializes tensor.extract on the gate-angle
// dataflow, so scalarization must run again here to fold the
- // extract_slice/collapse_shape chains it exposes (issue #2759).
+ // extract_slice/collapse_shape chains it exposes.
"scalarize-tensor-extracts",
"canonicalize",
"cse",
@@ -111,7 +111,7 @@ const PipelineList pipelineList{
// This pass is needed to avoid aliasing of the input buffer with the output buffer.
"mark-entry-point-args-non-writable",
// Value-number duplicate tensor computations before bufferization so they
- // do not each become a separate buffer (issue #2759).
+ // do not each become a separate buffer.
"cse",
"one-shot-bufferize",
// Remove dead memrefToTensorOp's
diff --git a/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
index 952188e4e8..2d6a28e9ab 100644
--- a/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
+++ b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
@@ -27,40 +27,17 @@
using namespace llvm;
using namespace mlir;
-// Tracing a Python circuit destroys its loop structure: N structurally identical
-// Trotter steps (or layers, folds, powers) arrive as N copies of the same op
-// sequence. Every later stage (bufferization, LLVM conversion, LLVM codegen)
-// amplifies each copy, which is the dominant source of IR blowup in issue #2759.
-//
-// This pass reconstructs those loops:
-//
-// 1. Canonical structural hashing: each op in a block is hashed over its name,
-// attributes, result types, and operand provenance, where operands defined
-// inside the block are identified by their *backward distance* in ops. The
-// hash sequence is shift-invariant, so identical iterations produce
-// identical hash subsequences regardless of position.
-//
-// 2. Tandem-repeat detection: maximal runs with H[i] == H[i+p] identify
-// candidate repeats w^k of period p (a linear-time shift-and-compare over
-// candidate periods harvested from same-hash gap statistics).
-//
-// 3. Semantic verification: consecutive windows must be isomorphic op-by-op,
-// with cross-window dataflow restricted to values *threaded* from the
-// directly preceding window at a fixed (position, result) slot — exactly the
-// shape of qubit values chained through a gate sequence — and all other
-// operands loop-invariant.
-//
-// 4. Materialization: the repeat is replaced by an scf.for whose iter_args are
-// the threaded slots, its body a clone of one window. A repeat of
-// multiplicity k shrinks that region k-fold before any downstream stage
-// sees it.
+// Tracing unrolls Python loops: N identical circuit segments (Trotter steps,
+// layers, folds) arrive as N copies of the same op sequence, and every later
+// stage amplifies each copy. This pass reconstructs the loops in four steps:
+// structural hashing of each op, tandem-repeat detection on the hash sequence
+// (maximal runs with H[i] == H[i+p]), semantic verification that consecutive
+// windows are isomorphic with cross-window dataflow limited to threaded values
+// (e.g. qubits) plus loop invariants, and replacement of the repeat with an
+// scf.for whose iter_args are the threaded values.
namespace {
-//===----------------------------------------------------------------------===//
-// Structural hashing
-//===----------------------------------------------------------------------===//
-
/// Ops that may participate in a rerolled window: single-block-region-free,
/// no successors, and not a terminator.
bool isRerollableOp(Operation *op)
@@ -69,18 +46,11 @@ bool isRerollableOp(Operation *op)
!op->hasTrait();
}
-/// Compute shift-invariant structural hashes for all ops of a block (excluding
-/// the terminator). Non-rerollable ops receive unique sentinel hashes so they
-/// can never be part of a repeat.
-///
-/// In-block defs are hashed with a constant tag plus the result number:
-/// hashing them by identity would make every window distinct (each window
-/// threads different SSA values), and hashing them by backward distance breaks
-/// on loop-invariant values (e.g. CSE-deduplicated angle computations) whose
-/// distance grows window over window. The resulting weaker discrimination is
-/// compensated by semantic verification of every candidate. Block arguments
-/// and out-of-block values are hashed by identity (loop-invariant by
-/// construction).
+/// Compute structural hashes for all ops of a block (excluding the
+/// terminator). Non-rerollable ops receive unique sentinel hashes. In-block
+/// operand defs are hashed by result number only (not identity, which would
+/// make every window distinct); the weak discrimination this leaves is
+/// compensated by semantic verification of every candidate.
void computeHashes(ArrayRef ops, const llvm::DenseMap &indexOf,
SmallVectorImpl &hashes)
{
@@ -110,10 +80,6 @@ void computeHashes(ArrayRef ops, const llvm::DenseMap hashes, unsigned minPeriod, unsigned maxP
}
}
-//===----------------------------------------------------------------------===//
-// Semantic verification
-//===----------------------------------------------------------------------===//
-
struct RerollPlan {
Candidate cand;
// Threaded slots: values flowing window -> next window, identified by
@@ -363,10 +325,7 @@ RerollPlan extendCandidate(ArrayRef ops,
return plan;
}
-//===----------------------------------------------------------------------===//
-// Materialization
-//===----------------------------------------------------------------------===//
-
+/// Replace the repeat with an scf.for and erase the original ops.
void materialize(ArrayRef ops, const RerollPlan &plan)
{
size_t start = plan.cand.start, p = plan.cand.period, count = plan.cand.count;
@@ -419,10 +378,6 @@ void materialize(ArrayRef ops, const RerollPlan &plan)
}
}
-//===----------------------------------------------------------------------===//
-// Driver
-//===----------------------------------------------------------------------===//
-
bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
{
SmallVector ops;
diff --git a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
index ef0c7665e9..90e202d6ea 100644
--- a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
+++ b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
@@ -25,24 +25,15 @@
using namespace llvm;
using namespace mlir;
-// Traced programs that compute quantum gate parameters from runtime inputs (e.g.
-// Trotterized Hamiltonians with runtime coefficients, see catalyst issue #2759)
-// produce long chains of *small* tensor ops whose only consumers are scalar
-// `tensor.extract` operations feeding gate parameters. Each such tensor survives
-// bufferization as an allocation plus copies, and each surviving `linalg.generic`
-// is later unrolled into loops, amplifying the IR by orders of magnitude.
-//
-// The patterns in this pass sink `tensor.extract` through the producers of such
-// tensors, computing the extracted element directly in scalar arithmetic:
-//
-// * extract(linalg.generic) -> inline the generic's scalar payload
-// * extract(tensor.collapse_shape) -> extract from the source with expanded indices
-// * extract(tensor.extract_slice) -> extract from the source with offset indices
-//
-// Applied to a fixpoint, extraction sinks to the leaves of the dataflow, the
-// intermediate tensors become dead, and the scalar chains remain (deduplicated by
-// a follow-up CSE). To avoid code growth on large tensors, generic-payload
-// inlining is restricted to small statically-shaped results.
+// Gate parameters computed from runtime inputs (e.g. Trotterization with
+// runtime Hamiltonian coefficients) arrive as long chains of small tensor ops
+// consumed only by scalar `tensor.extract` operations; each such tensor
+// survives bufferization as an allocation plus copies. This pass sinks
+// `tensor.extract` through `linalg.generic` (inlining the scalar payload),
+// `tensor.collapse_shape`, and `tensor.extract_slice`, so the extracted
+// element is computed in scalar arithmetic and the tensors become dead.
+// Payload inlining is limited to small statically-shaped results to bound
+// code growth.
namespace {
From e01ebbbf82b4e5ab82905316cd9974b8709a9d3f Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Sun, 12 Jul 2026 03:23:41 -0400
Subject: [PATCH 3/9] Fix BufferizationStage IR amplification for
runtime-coefficient Hamiltonians
Add scalarize-tensor-extracts and reroll-loops passes and re-order the
default pipeline so runtime-coefficient Trotterization no longer amplifies
IR through bufferization (9.9x -> ~1.0x on the H2 QPE benchmark).
Fixes #2759
---
doc/releases/changelog-dev.md | 21 +--
.../pytest/test_trotter_runtime_coeffs.py | 26 +--
mlir/include/Catalyst/Transforms/Passes.td | 6 +-
.../DefaultPipelines/DefaultPipelines.h | 9 +-
.../Catalyst/Transforms/RerollLoopsPass.cpp | 150 +++++++++++++-----
.../ScalarizeTensorExtractsPass.cpp | 23 ++-
mlir/test/Catalyst/RerollLoopsTest.mlir | 30 ++--
7 files changed, 163 insertions(+), 102 deletions(-)
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index ead45f46d7..2f8e7b157b 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -161,20 +161,13 @@
* IR amplification for circuits whose gate parameters are computed from runtime values
(e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
- Three changes work together in the default pipeline:
- the new `scalarize-tensor-extracts` pass sinks scalar `tensor.extract` operations through
- `tensor.extract_slice`, `tensor.collapse_shape`, and small elementwise `linalg.generic`
- producers so gate-angle dataflow becomes pure scalar arithmetic instead of thousands of
- tiny tensors that each survive bufferization as an allocation and copies;
- `linalg-fuse-elementwise-ops` plus additional `cse` applications (before and after
- `one-shot-bufferize`) remove the duplicate angle computations produced by tracing; and
- the new `reroll-loops` pass reconstructs the loops that Python tracing unrolled, by
- detecting tandem repeats of structurally isomorphic operation windows (via structural
- hashing), verifying that cross-window dataflow is limited to threaded SSA values (such as
- qubit values) plus loop-invariant values, and replacing each repeat with an `scf.for`.
- For a representative Trotterized QPE workload with runtime coefficients, gate-op volume
- after HLO lowering drops about 5x and downstream IR, compile time, and peak memory drop
- accordingly.
+ The new `scalarize-tensor-extracts` pass turns gate-angle dataflow into scalar arithmetic
+ instead of thousands of small tensors that each survive bufferization as an allocation,
+ the new `reroll-loops` pass reconstructs the loops that tracing unrolled by rewriting
+ repeated op sequences as `scf.for` loops, and the default pipeline now runs elementwise
+ fusion. On a Trotterized QPE workload with runtime coefficients, compile time, peak
+ memory, and final IR size all drop by large factors.
+ [(#3013)](https://github.com/PennyLaneAI/catalyst/pull/3013)
Breaking changes 💔
diff --git a/frontend/test/pytest/test_trotter_runtime_coeffs.py b/frontend/test/pytest/test_trotter_runtime_coeffs.py
index b71720589b..d1efeed4a6 100644
--- a/frontend/test/pytest/test_trotter_runtime_coeffs.py
+++ b/frontend/test/pytest/test_trotter_runtime_coeffs.py
@@ -89,23 +89,27 @@ def test_runtime_matches_fixed(self):
assert np.allclose(np.asarray(dyn), np.asarray(fixed), atol=1e-9)
def test_reroll_recovers_loops(self):
- """The compiled module must contain scf.for loops recovered from the
- unrolled Trotter steps (guards against silent regression of
- reroll-loops in the default pipeline)."""
+ """The IR after HLO lowering must contain scf.for loops recovered from
+ the unrolled Trotter steps, and fewer gate ops than the unrolled
+ circuit (guards against silent regression of reroll-loops in the
+ default pipeline)."""
+ from catalyst.debug import get_compilation_stage
+
dev = qml.device("lightning.qubit", wires=N_QUBITS + N_EST)
coeffs = jnp.array(COEFFS)
compiled = qjit(make_qpe(dev, runtime=True), keep_intermediate=True)
compiled(coeffs)
try:
- workspace = str(compiled.workspace)
- import glob
- import os
-
- hlo_files = glob.glob(os.path.join(workspace, "*HLOLowering*.mlir"))
- assert hlo_files, "no post-HLO snapshot written"
- content = open(hlo_files[0], encoding="utf-8").read()
- assert "scf.for" in content, "reroll-loops did not fire"
+ traced = get_compilation_stage(compiled, "QuantumCompilationStage")
+ lowered = get_compilation_stage(compiled, "HLOLoweringStage")
+ assert "scf.for" in lowered, "reroll-loops did not fire"
+ unrolled_gates = traced.count("quantum.custom")
+ rerolled_gates = lowered.count("quantum.custom")
+ assert rerolled_gates < unrolled_gates / 2, (
+ f"expected reroll to shrink gate volume by >2x, got "
+ f"{unrolled_gates} -> {rerolled_gates}"
+ )
finally:
compiled.workspace.cleanup()
diff --git a/mlir/include/Catalyst/Transforms/Passes.td b/mlir/include/Catalyst/Transforms/Passes.td
index 37ab0667fd..c1c52e2735 100644
--- a/mlir/include/Catalyst/Transforms/Passes.td
+++ b/mlir/include/Catalyst/Transforms/Passes.td
@@ -365,9 +365,11 @@ def RerollLoopsPass : Pass<"reroll-loops"> {
/*C++ var name=*/"minSavings",
/*CLI arg name=*/"min-savings",
/*type=*/"unsigned",
- /*default=*/"8",
+ /*default=*/"32",
/*description=*/
- "Minimum number of ops a reroll must eliminate to be applied."
+ "Minimum number of ops a reroll must eliminate to be applied. The "
+ "default is deliberately conservative so that only substantial "
+ "repeats (e.g. unrolled Trotter steps) are converted to loops."
>
];
}
diff --git a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
index 7a308085f1..7904824045 100644
--- a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
+++ b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
@@ -78,7 +78,6 @@ const PipelineList pipelineList{
"scalarize-tensor-extracts",
"func.func(linalg-fuse-elementwise-ops)",
"canonicalize",
- "cse",
// Reconstruct the loops that tracing unrolled (e.g. Trotter steps); a
// repeat of multiplicity k shrinks its region k-fold before the
// bufferization and LLVM stages amplify it.
@@ -91,7 +90,6 @@ const PipelineList pipelineList{
// extract_slice/collapse_shape chains it exposes.
"scalarize-tensor-extracts",
"canonicalize",
- "cse",
"symbol-dce"}},
{"gradient-lowering-stage",
{"annotate-invalid-gradient-functions",
@@ -110,9 +108,6 @@ const PipelineList pipelineList{
*/
// This pass is needed to avoid aliasing of the input buffer with the output buffer.
"mark-entry-point-args-non-writable",
- // Value-number duplicate tensor computations before bufferization so they
- // do not each become a separate buffer.
- "cse",
"one-shot-bufferize",
// Remove dead memrefToTensorOp's
"canonicalize",
@@ -132,7 +127,9 @@ const PipelineList pipelineList{
// Must be after convert-bufferization-to-memref.
// Otherwise, there are issues in the lowering of dynamic tensors.
"canonicalize",
- "cse",
+ /* [DISABLED PASS]
+ * "cse",
+ */
"cp-global-memref"}},
{"llvm-dialect-lowering-stage",
{"qnode-to-async-lowering",
diff --git a/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
index 2d6a28e9ab..b60e5265bf 100644
--- a/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
+++ b/mlir/lib/Catalyst/Transforms/RerollLoopsPass.cpp
@@ -17,7 +17,6 @@
#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/Hashing.h"
#include "llvm/Support/Debug.h"
-
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/IR/Builders.h"
@@ -38,6 +37,26 @@ using namespace mlir;
namespace {
+/// Verifying a candidate period is linear in the block, so bound how many
+/// distinct periods (most-promising first) are tried per block.
+constexpr size_t kMaxPeriodsPerBlock = 16;
+
+/// Minimum number of windows for a repeat to become a loop. Two windows are
+/// cheaper as straight-line code and give extension no signal to distinguish
+/// coincidental hash matches from real repetition.
+constexpr size_t kMinWindows = 3;
+
+/// Rerolling creates new blocks (loop bodies) that may themselves contain
+/// repeats (nested loops), so the driver iterates; each round strictly
+/// shrinks the IR, and in practice one nesting level per round suffices.
+constexpr unsigned kMaxRounds = 4;
+
+/// Number of rounds of operand-hash mixing (see refineHashes). Deep enough to
+/// tell apart structurally similar ops at different positions of a window
+/// (e.g. the same gate applied to different qubits of a chain), shallow enough
+/// that an op's hash rarely depends on ops more than a window away.
+constexpr unsigned kHashDepth = 3;
+
/// Ops that may participate in a rerolled window: single-block-region-free,
/// no successors, and not a terminator.
bool isRerollableOp(Operation *op)
@@ -47,10 +66,11 @@ bool isRerollableOp(Operation *op)
}
/// Compute structural hashes for all ops of a block (excluding the
-/// terminator). Non-rerollable ops receive unique sentinel hashes. In-block
-/// operand defs are hashed by result number only (not identity, which would
-/// make every window distinct); the weak discrimination this leaves is
-/// compensated by semantic verification of every candidate.
+/// terminator). Non-rerollable ops receive unique sentinel hashes. The hash
+/// covers the op name, attributes, and result types; in-block operand defs
+/// contribute only their result number (hashing them by identity would make
+/// every window distinct), out-of-block values their identity. The weak
+/// discrimination this leaves is compensated by semantic verification.
void computeHashes(ArrayRef ops, const llvm::DenseMap &indexOf,
SmallVectorImpl &hashes)
{
@@ -60,8 +80,8 @@ void computeHashes(ArrayRef ops, const llvm::DenseMapgetName().getTypeID(),
- op->getAttrDictionary().getAsOpaquePointer());
+ hash_code h =
+ hash_combine(op->getName().getTypeID(), op->getAttrDictionary().getAsOpaquePointer());
for (Type t : op->getResultTypes()) {
h = hash_combine(h, t.getAsOpaquePointer());
}
@@ -80,6 +100,35 @@ void computeHashes(ArrayRef ops, const llvm::DenseMap ops, const llvm::DenseMap &indexOf,
+ SmallVectorImpl &hashes)
+{
+ SmallVector prev(hashes.begin(), hashes.end());
+ for (const auto &[i, op] : llvm::enumerate(ops)) {
+ if (!isRerollableOp(op)) {
+ continue;
+ }
+ hash_code h = prev[i];
+ for (Value v : op->getOperands()) {
+ Operation *def = v.getDefiningOp();
+ auto it = def ? indexOf.find(def) : indexOf.end();
+ if (it != indexOf.end()) {
+ h = hash_combine(h, prev[it->second]);
+ }
+ }
+ hashes[i] = h;
+ }
+}
+
struct Candidate {
size_t start; // index of the first op of the first window
size_t period; // ops per window
@@ -107,16 +156,16 @@ void findCandidates(ArrayRef hashes, unsigned minPeriod, unsigned maxP
SmallVector> periods(gapWeight.begin(), gapWeight.end());
// Favor periods that explain the most ops.
- llvm::sort(periods, [](auto &a, auto &b) {
- return a.first * a.second > b.first * b.second;
- });
+ llvm::sort(periods, [](auto &a, auto &b) { return a.first * a.second > b.first * b.second; });
if (periods.size() > maxPeriods) {
periods.resize(maxPeriods);
}
for (auto &[p, weight] : periods) {
- // Require the period to repeat enough times to be worth a loop.
- if (weight < 2 * p) {
+ // A repeat of period p with >= kMinWindows windows produces at
+ // least (kMinWindows - 1) * p same-hash pairs at gap p; anything below
+ // that cannot yield an acceptable candidate.
+ if (weight < (kMinWindows - 1) * p) {
continue;
}
size_t runStart = SIZE_MAX;
@@ -128,7 +177,7 @@ void findCandidates(ArrayRef hashes, unsigned minPeriod, unsigned maxP
if (!match && runStart != SIZE_MAX) {
size_t runLen = i - runStart;
size_t count = runLen / p + 1;
- if (count >= 3) {
+ if (count >= kMinWindows) {
out.push_back({runStart, p, count});
}
runStart = SIZE_MAX;
@@ -182,8 +231,7 @@ std::optional verifyCandidate(ArrayRef ops,
if (!isRerollableOp(a) || !isRerollableOp(b)) {
return std::nullopt;
}
- if (a->getName() != b->getName() ||
- a->getAttrDictionary() != b->getAttrDictionary() ||
+ if (a->getName() != b->getName() || a->getAttrDictionary() != b->getAttrDictionary() ||
a->getResultTypes() != b->getResultTypes() ||
a->getNumOperands() != b->getNumOperands()) {
return std::nullopt;
@@ -276,8 +324,7 @@ std::optional verifyCandidate(ArrayRef ops,
size_t ui = (uit != indexOf.end()) ? uit->second : SIZE_MAX;
bool inOwnWindow =
ui != SIZE_MAX && ui >= start + w * p && ui < start + (w + 1) * p;
- bool inNextWindow = !isLast && ui != SIZE_MAX &&
- ui >= start + (w + 1) * p &&
+ bool inNextWindow = !isLast && ui != SIZE_MAX && ui >= start + (w + 1) * p &&
ui < start + (w + 2) * p;
if (inOwnWindow || inNextWindow) {
continue;
@@ -299,28 +346,44 @@ std::optional verifyCandidate(ArrayRef ops,
/// Try to extend a verified candidate by whole windows to the left/right; the
/// hash sequence misses the first window (its cross-window references point at
-/// the prologue, at different distances), so this recovers it.
+/// the prologue, at different distances), so this recovers it. Each
+/// verification is linear in the candidate size, so the number of windows
+/// added per attempt grows geometrically (and resets on failure) to keep the
+/// total cost O(size * log(windows)) rather than quadratic.
RerollPlan extendCandidate(ArrayRef ops,
const llvm::DenseMap &indexOf, RerollPlan plan)
{
- while (plan.cand.start >= plan.cand.period) {
+ size_t step = 1;
+ while (plan.cand.start >= plan.cand.period * step) {
Candidate c = plan.cand;
- c.start -= c.period;
- c.count += 1;
- auto extended = verifyCandidate(ops, indexOf, c);
- if (!extended) {
+ c.start -= c.period * step;
+ c.count += step;
+ if (auto extended = verifyCandidate(ops, indexOf, c)) {
+ plan = *extended;
+ step *= 2;
+ }
+ else if (step == 1) {
break;
}
- plan = *extended;
+ else {
+ step = 1;
+ }
}
+ step = 1;
while (true) {
Candidate c = plan.cand;
- c.count += 1;
- auto extended = verifyCandidate(ops, indexOf, c);
- if (!extended) {
+ c.count += step;
+ if (c.start + c.count * c.period <= ops.size()) {
+ if (auto extended = verifyCandidate(ops, indexOf, c)) {
+ plan = *extended;
+ step *= 2;
+ continue;
+ }
+ }
+ if (step == 1) {
break;
}
- plan = *extended;
+ step = 1;
}
return plan;
}
@@ -371,9 +434,11 @@ void materialize(ArrayRef ops, const RerollPlan &plan)
lastOp->getResult(resNo).replaceAllUsesWith(forOp.getResult(slotIdx));
}
- // Erase the original ops, last first (uses before defs).
+ // Erase the original ops, last first (uses before defs). Verification
+ // guarantees no surviving uses; erase() asserts use_empty, so a
+ // verification bug fails loudly here instead of producing invalid IR.
for (size_t i = start + count * p; i-- > start;) {
- ops[i]->dropAllUses();
+ assert(ops[i]->use_empty() && "rerolled op still has uses; verification is unsound");
ops[i]->erase();
}
}
@@ -393,11 +458,18 @@ bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
indexOf[op] = i;
}
+ // Harvest candidates at every hash refinement depth: shallow hashes see
+ // repeats whose windows contain few distinct op kinds only as noise, deep
+ // hashes lose the first windows of a run to prologue ancestry. Duplicated
+ // candidates are cheap (verification dedups via the overlap check).
SmallVector hashes;
computeHashes(ops, indexOf, hashes);
-
SmallVector candidates;
- findCandidates(hashes, minPeriod, /*maxPeriods=*/16, candidates);
+ findCandidates(hashes, minPeriod, kMaxPeriodsPerBlock, candidates);
+ for (unsigned depth = 0; depth < kHashDepth; ++depth) {
+ refineHashes(ops, indexOf, hashes);
+ findCandidates(hashes, minPeriod, kMaxPeriodsPerBlock, candidates);
+ }
// Verify, extend, and pick non-overlapping plans greedily by savings.
SmallVector plans;
@@ -419,8 +491,8 @@ bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
SmallVector accepted;
for (const RerollPlan &plan : plans) {
size_t s = plan.cand.start, e = s + plan.cand.count * plan.cand.period;
- bool overlaps = llvm::any_of(
- used, [&](auto range) { return s < range.second && range.first < e; });
+ bool overlaps =
+ llvm::any_of(used, [&](auto range) { return s < range.second && range.first < e; });
if (!overlaps) {
used.push_back({s, e});
accepted.push_back(&plan);
@@ -434,8 +506,7 @@ bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
});
for (const RerollPlan *plan : accepted) {
LLVM_DEBUG(dbgs() << "rerolling: start=" << plan->cand.start
- << " period=" << plan->cand.period
- << " count=" << plan->cand.count
+ << " period=" << plan->cand.period << " count=" << plan->cand.count
<< " slots=" << plan->slots.size() << "\n");
materialize(ops, *plan);
}
@@ -446,6 +517,9 @@ bool processBlock(Block &block, unsigned minPeriod, unsigned minSavings)
namespace catalyst {
+// GEN_PASS_DECL is needed in addition to GEN_PASS_DEF to declare the
+// RerollLoopsPassOptions struct that the generated base class references
+// (only passes with options need this).
#define GEN_PASS_DECL_REROLLLOOPSPASS
#define GEN_PASS_DEF_REROLLLOOPSPASS
#include "Catalyst/Transforms/Passes.h.inc"
@@ -455,11 +529,9 @@ struct RerollLoopsPass : public impl::RerollLoopsPassBase {
void runOnOperation() override
{
- // Iterate to a fixpoint (bounded): rerolling creates new blocks (loop
- // bodies) that may contain further repeats, e.g. nested loops.
bool changed = true;
unsigned rounds = 0;
- while (changed && rounds++ < 4) {
+ while (changed && rounds++ < kMaxRounds) {
changed = false;
SmallVector blocks;
getOperation()->walk([&](Block *block) { blocks.push_back(block); });
diff --git a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
index 90e202d6ea..858ab70026 100644
--- a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
+++ b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
@@ -38,7 +38,9 @@ using namespace mlir;
namespace {
/// Upper bound on the number of elements of a tensor whose producer payload we
-/// are willing to clone per extraction site.
+/// are willing to clone per extraction site. 16 covers the 4x4 matrices of
+/// two-qubit gates, the largest tensors in gate-parameter dataflow, while
+/// keeping the worst-case code growth per extract small.
constexpr int64_t kMaxScalarizedElements = 16;
/// Fold tensor.extract(linalg.generic) by inlining the generic's scalar payload
@@ -114,8 +116,8 @@ struct ExtractOfGeneric : public OpRewritePattern {
inputIndices.push_back(iterIndices[dimExpr.getPosition()]);
}
else if (auto constExpr = dyn_cast(expr)) {
- inputIndices.push_back(arith::ConstantIndexOp::create(
- rewriter, loc, constExpr.getValue()));
+ inputIndices.push_back(
+ arith::ConstantIndexOp::create(rewriter, loc, constExpr.getValue()));
}
else {
return failure();
@@ -188,14 +190,12 @@ struct ExtractOfCollapseShape : public OpRewritePattern {
}
}
for (int64_t srcDim : group) {
- srcIndices[srcDim] = (srcDim == nonUnitDim)
- ? extractOp.getIndices()[groupIdx]
- : getZero();
+ srcIndices[srcDim] =
+ (srcDim == nonUnitDim) ? extractOp.getIndices()[groupIdx] : getZero();
}
}
- rewriter.replaceOpWithNewOp(extractOp, collapseOp.getSrc(),
- srcIndices);
+ rewriter.replaceOpWithNewOp(extractOp, collapseOp.getSrc(), srcIndices);
return success();
}
};
@@ -224,8 +224,8 @@ struct ExtractOfExtractSlice : public OpRewritePattern {
if (auto val = dyn_cast(ofr)) {
return val;
}
- return arith::ConstantIndexOp::create(
- rewriter, loc, cast(cast(ofr)).getInt());
+ return arith::ConstantIndexOp::create(rewriter, loc,
+ cast(cast(ofr)).getInt());
};
SmallVector srcIndices;
@@ -242,8 +242,7 @@ struct ExtractOfExtractSlice : public OpRewritePattern {
srcIndices.push_back(arith::AddIOp::create(rewriter, loc, offset, scaled));
}
- rewriter.replaceOpWithNewOp(extractOp, sliceOp.getSource(),
- srcIndices);
+ rewriter.replaceOpWithNewOp(extractOp, sliceOp.getSource(), srcIndices);
return success();
}
};
diff --git a/mlir/test/Catalyst/RerollLoopsTest.mlir b/mlir/test/Catalyst/RerollLoopsTest.mlir
index 28a59dd007..5fe1ee9ad0 100644
--- a/mlir/test/Catalyst/RerollLoopsTest.mlir
+++ b/mlir/test/Catalyst/RerollLoopsTest.mlir
@@ -15,17 +15,14 @@
// RUN: quantum-opt --reroll-loops="min-period=2 min-savings=4" --split-input-file %s | FileCheck %s
// A scalar chain of alternating ops rerolls into an scf.for threading one
-// value. The first addf seeds the loop (its operand pattern differs), the
-// trailing mulf stays as an epilogue.
+// value through all eight iterations.
// CHECK-LABEL: @scalar_chain
-// CHECK: %[[SEED:.+]] = arith.addf
-// CHECK: %[[FOR:.+]] = scf.for {{.*}} iter_args(%[[IT:.+]] = %[[SEED]]) -> (f64)
-// CHECK: %[[M:.+]] = arith.mulf %[[IT]],
-// CHECK: %[[A:.+]] = arith.addf %[[M]],
-// CHECK: scf.yield %[[A]] : f64
-// CHECK: %[[EPI:.+]] = arith.mulf %[[FOR]],
-// CHECK: return %[[EPI]]
+// CHECK: %[[FOR:.+]] = scf.for {{.*}} iter_args(%[[IT:.+]] = %arg0) -> (f64)
+// CHECK: %[[A:.+]] = arith.addf %[[IT]],
+// CHECK: %[[M:.+]] = arith.mulf %[[A]],
+// CHECK: scf.yield %[[M]] : f64
+// CHECK: return %[[FOR]]
func.func @scalar_chain(%arg0: f64, %c: f64) -> f64 {
%0 = arith.addf %arg0, %c : f64
%1 = arith.mulf %0, %c : f64
@@ -49,19 +46,16 @@ func.func @scalar_chain(%arg0: f64, %c: f64) -> f64 {
// -----
// A repeated gate sequence threading two qubits rerolls with both qubit
-// values as iter_args; the rotation angle is loop-invariant. The pass may pick
-// any rotation of the repeated window (here the run starts at the first CNOT),
-// leaving a prologue/epilogue outside the loop.
+// values as iter_args; the rotation angle is loop-invariant.
// CHECK-LABEL: @gate_sequence
// CHECK: quantum.alloc
// CHECK: %[[FOR:.+]]:2 = scf.for {{.*}} iter_args(%[[Q0:.+]] = %{{.+}}, %[[Q1:.+]] = %{{.+}}) -> (!quantum.bit, !quantum.bit)
-// CHECK: %[[CNOT:.+]]:2 = quantum.custom "CNOT"() %[[Q0]], %[[Q1]]
-// CHECK: %[[H:.+]] = quantum.custom "Hadamard"() %[[CNOT]]#0
-// CHECK: %[[RZ:.+]] = quantum.custom "RZ"(%{{.+}}) %[[CNOT]]#1
-// CHECK: scf.yield %[[H]], %[[RZ]]
-// CHECK: %[[LAST:.+]]:2 = quantum.custom "CNOT"() %[[FOR]]#0, %[[FOR]]#1
-// CHECK: quantum.insert %{{.+}}[ 0], %[[LAST]]#0
+// CHECK: %[[H:.+]] = quantum.custom "Hadamard"() %[[Q0]]
+// CHECK: %[[RZ:.+]] = quantum.custom "RZ"(%{{.+}}) %[[Q1]]
+// CHECK: %[[CNOT:.+]]:2 = quantum.custom "CNOT"() %[[H]], %[[RZ]]
+// CHECK: scf.yield %[[CNOT]]#0, %[[CNOT]]#1
+// CHECK: quantum.insert %{{.+}}[ 0], %[[FOR]]#0
func.func @gate_sequence(%theta: f64) -> !quantum.reg {
%r0 = quantum.alloc( 2) : !quantum.reg
%q0 = quantum.extract %r0[ 0] : !quantum.reg -> !quantum.bit
From a4612115412c0f05d1a4930e825b97dfc406f025 Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Sun, 12 Jul 2026 03:43:10 -0400
Subject: [PATCH 4/9] update certain things
---
.../pytest/test_trotter_runtime_coeffs.py | 9 +++-
.../DefaultPipelines/DefaultPipelines.h | 7 ++-
.../ScalarizeTensorExtractsPass.cpp | 51 +++++++++++++------
3 files changed, 49 insertions(+), 18 deletions(-)
diff --git a/frontend/test/pytest/test_trotter_runtime_coeffs.py b/frontend/test/pytest/test_trotter_runtime_coeffs.py
index d1efeed4a6..6b0bd79b43 100644
--- a/frontend/test/pytest/test_trotter_runtime_coeffs.py
+++ b/frontend/test/pytest/test_trotter_runtime_coeffs.py
@@ -17,6 +17,8 @@
in the default pipeline must preserve numerics for Trotterized workloads with
runtime coefficients."""
+import re
+
import numpy as np
import pennylane as qml
import pytest
@@ -103,7 +105,12 @@ def test_reroll_recovers_loops(self):
try:
traced = get_compilation_stage(compiled, "QuantumCompilationStage")
lowered = get_compilation_stage(compiled, "HLOLoweringStage")
- assert "scf.for" in lowered, "reroll-loops did not fire"
+ # Rerolled Trotter steps are scf.for loops threading qubit values
+ # through iter_args.
+ qubit_loops = re.findall(
+ r"scf\.for .*iter_args\([^)]*\).*->.*!quantum\.bit", lowered
+ )
+ assert qubit_loops, "reroll-loops did not produce qubit-threading loops"
unrolled_gates = traced.count("quantum.custom")
rerolled_gates = lowered.count("quantum.custom")
assert rerolled_gates < unrolled_gates / 2, (
diff --git a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
index 7904824045..7e4a7b40b3 100644
--- a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
+++ b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
@@ -80,7 +80,12 @@ const PipelineList pipelineList{
"canonicalize",
// Reconstruct the loops that tracing unrolled (e.g. Trotter steps); a
// repeat of multiplicity k shrinks its region k-fold before the
- // bufferization and LLVM stages amplify it.
+ // bufferization and LLVM stages amplify it. No cse is needed before
+ // this pass computations duplicated per iteration sit inside each
+ // repeated window and reroll as part of it, and constants (the only
+ // cross-window operands that must be identical SSA values) are already
+ // uniqued by the preceding canonicalize. An ablation with cse here
+ // showed no change in rerolling or output size.
"reroll-loops",
"func.func(linalg-detensorize{aggressive-mode})",
"detensorize-scf",
diff --git a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
index 858ab70026..e60aca1a96 100644
--- a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
+++ b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
@@ -94,6 +94,20 @@ struct ExtractOfGeneric : public OpRewritePattern {
}
}
+ // Validate every input indexing map up front: patterns must not mutate
+ // the IR on the failure path, and ops are created per input below.
+ for (OpOperand *inOperand : genericOp.getDpsInputOperands()) {
+ if (!isa(inOperand->get().getType())) {
+ continue;
+ }
+ AffineMap inputMap = genericOp.getMatchingIndexingMap(inOperand);
+ for (AffineExpr expr : inputMap.getResults()) {
+ if (!isa(expr)) {
+ return failure();
+ }
+ }
+ }
+
Location loc = extractOp.getLoc();
SmallVector iterIndices(extractOp.getIndices());
@@ -115,13 +129,11 @@ struct ExtractOfGeneric : public OpRewritePattern {
if (auto dimExpr = dyn_cast(expr)) {
inputIndices.push_back(iterIndices[dimExpr.getPosition()]);
}
- else if (auto constExpr = dyn_cast(expr)) {
+ else {
+ auto constExpr = cast(expr);
inputIndices.push_back(
arith::ConstantIndexOp::create(rewriter, loc, constExpr.getValue()));
}
- else {
- return failure();
- }
}
Value scalar = tensor::ExtractOp::create(rewriter, loc, input, inputIndices);
mapping.map(blockArg, scalar);
@@ -161,6 +173,22 @@ struct ExtractOfCollapseShape : public OpRewritePattern {
return failure();
}
+ // Validate every reassociation group up front (at most one non-unit
+ // source dimension per group): patterns must not mutate the IR on the
+ // failure path, and index constants are created per group below.
+ SmallVector groups = collapseOp.getReassociationIndices();
+ SmallVector nonUnitDims(groups.size(), -1);
+ for (const auto &[groupIdx, group] : llvm::enumerate(groups)) {
+ for (int64_t srcDim : group) {
+ if (srcType.getDimSize(srcDim) != 1) {
+ if (nonUnitDims[groupIdx] != -1) {
+ return failure(); // true merge of two non-unit dims
+ }
+ nonUnitDims[groupIdx] = srcDim;
+ }
+ }
+ }
+
Location loc = extractOp.getLoc();
SmallVector srcIndices(srcType.getRank());
@@ -172,7 +200,6 @@ struct ExtractOfCollapseShape : public OpRewritePattern {
return zero;
};
- SmallVector groups = collapseOp.getReassociationIndices();
// A rank-0 result means every source dimension is a unit dimension.
if (groups.empty()) {
for (int64_t dim = 0; dim < srcType.getRank(); ++dim) {
@@ -180,18 +207,10 @@ struct ExtractOfCollapseShape : public OpRewritePattern {
}
}
for (const auto &[groupIdx, group] : llvm::enumerate(groups)) {
- int64_t nonUnitDim = -1;
- for (int64_t srcDim : group) {
- if (srcType.getDimSize(srcDim) != 1) {
- if (nonUnitDim != -1) {
- return failure(); // true merge of two non-unit dims
- }
- nonUnitDim = srcDim;
- }
- }
for (int64_t srcDim : group) {
- srcIndices[srcDim] =
- (srcDim == nonUnitDim) ? extractOp.getIndices()[groupIdx] : getZero();
+ srcIndices[srcDim] = (srcDim == nonUnitDims[groupIdx])
+ ? extractOp.getIndices()[groupIdx]
+ : getZero();
}
}
From 5b2e9a4de059602500431936475e7f6461894cd8 Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Thu, 16 Jul 2026 15:28:59 -0400
Subject: [PATCH 5/9] Remove scalarize-tensor-extracts; this branch carries
only the reroll-loops pass (split from #3013)
Co-authored-by: Cursor
---
doc/releases/changelog-dev.md | 16 +-
frontend/test/lit/test_detensorize.py | 8 +-
...runtime_coeffs.py => test_reroll_loops.py} | 13 +-
mlir/include/Catalyst/Transforms/Passes.td | 24 --
.../DefaultPipelines/DefaultPipelines.h | 21 +-
mlir/lib/Catalyst/Transforms/CMakeLists.txt | 1 -
.../ScalarizeTensorExtractsPass.cpp | 292 ------------------
.../Catalyst/ScalarizeTensorExtractsTest.mlir | 146 ---------
8 files changed, 21 insertions(+), 500 deletions(-)
rename frontend/test/pytest/{test_trotter_runtime_coeffs.py => test_reroll_loops.py} (90%)
delete mode 100644 mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
delete mode 100644 mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index ab7f8695bb..9d513614ef 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -166,15 +166,13 @@
* Added ``CZ`` support to ``to-ppr`` pass.
[(#3009)](https://github.com/PennyLaneAI/catalyst/pull/3009)
-* IR amplification for circuits whose gate parameters are computed from runtime values
- (e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
- The new `scalarize-tensor-extracts` pass turns gate-angle dataflow into scalar arithmetic
- instead of thousands of small tensors that each survive bufferization as an allocation,
- the new `reroll-loops` pass reconstructs the loops that tracing unrolled by rewriting
- repeated op sequences as `scf.for` loops, and the default pipeline now runs elementwise
- fusion. On a Trotterized QPE workload with runtime coefficients, compile time, peak
- memory, and final IR size all drop by large factors.
- [(#3013)](https://github.com/PennyLaneAI/catalyst/pull/3013)
+* A new `reroll-loops` pass reconstructs the loops that tracing unrolled by detecting
+ repeated op sequences (e.g. Trotter steps, layers, folds) via structural hashing and
+ rewriting them as `scf.for` loops. A repeat of multiplicity k shrinks that IR region
+ k-fold before the bufferization and LLVM stages amplify it, reducing compile time,
+ peak memory, and IR size on workloads with many structurally identical circuit
+ segments. The pass runs by default in the HLO lowering stage.
+ [(#XXXX)](https://github.com/PennyLaneAI/catalyst/pull/XXXX)
Breaking changes 💔
diff --git a/frontend/test/lit/test_detensorize.py b/frontend/test/lit/test_detensorize.py
index 40cf98509b..cc2a2295cd 100644
--- a/frontend/test/lit/test_detensorize.py
+++ b/frontend/test/lit/test_detensorize.py
@@ -45,15 +45,15 @@ def f_with_cond(a, b):
# CHECK-NOT: linalg.generic
a2 = a + a
if a2 > b:
- # CHECK-DAG: arith.subf
+ # CHECK: arith.subf
a = a - 2.0
- # CHECK-DAG: arith.mulf
+ # CHECK: arith.mulf
b = b * 2.0
c = a + b
else:
- # CHECK-DAG: arith.mulf
+ # CHECK: arith.mulf
a = a * 2.0
- # CHECK-DAG: arith.subf
+ # CHECK: arith.subf
b = b - 2.0
c = a + b
return c * 2.0
diff --git a/frontend/test/pytest/test_trotter_runtime_coeffs.py b/frontend/test/pytest/test_reroll_loops.py
similarity index 90%
rename from frontend/test/pytest/test_trotter_runtime_coeffs.py
rename to frontend/test/pytest/test_reroll_loops.py
index 6b0bd79b43..f1c35672d7 100644
--- a/frontend/test/pytest/test_trotter_runtime_coeffs.py
+++ b/frontend/test/pytest/test_reroll_loops.py
@@ -12,10 +12,9 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""Integration tests for the IR-amplification fixes for runtime-coefficient
-Hamiltonians: scalarize-tensor-extracts, elementwise fusion, and reroll-loops
-in the default pipeline must preserve numerics for Trotterized workloads with
-runtime coefficients."""
+"""Integration tests for the reroll-loops pass in the default pipeline:
+rerolling the unrolled Trotter steps of a QPE workload must preserve numerics
+and must actually reconstruct scf.for loops."""
import re
@@ -75,9 +74,9 @@ def qpe_circuit(coeffs):
return qpe_circuit
-class TestRuntimeCoefficientTrotter:
- """Numerical equivalence of runtime- and fixed-coefficient Trotterization
- through the default pipeline (which scalarizes, fuses, and rerolls)."""
+class TestRerollLoops:
+ """Numerical equivalence and loop reconstruction for Trotterized workloads
+ through the default pipeline (which rerolls unrolled repeats)."""
def test_runtime_matches_fixed(self):
"""qml.dot with traced coefficients must produce the same distribution
diff --git a/mlir/include/Catalyst/Transforms/Passes.td b/mlir/include/Catalyst/Transforms/Passes.td
index c1c52e2735..6196e4eff7 100644
--- a/mlir/include/Catalyst/Transforms/Passes.td
+++ b/mlir/include/Catalyst/Transforms/Passes.td
@@ -374,30 +374,6 @@ def RerollLoopsPass : Pass<"reroll-loops"> {
];
}
-def ScalarizeTensorExtractsPass : Pass<"scalarize-tensor-extracts"> {
- let summary = "Sink scalar tensor.extract ops through small-tensor producers.";
- let description = [{
- Programs that compute quantum gate parameters from runtime inputs (e.g.
- Trotterization with runtime Hamiltonian coefficients) produce long chains of
- small tensor operations whose only consumers are scalar `tensor.extract`
- operations. Each intermediate tensor survives bufferization as an allocation
- with copies and each `linalg.generic` is later unrolled into loops, causing
- severe IR amplification.
-
- This pass folds `tensor.extract` through elementwise `linalg.generic` (by
- inlining the scalar payload), `tensor.collapse_shape`, and
- `tensor.extract_slice`, so the extracted element is computed directly in
- scalar arithmetic and the intermediate tensors become dead. Payload inlining
- is limited to small statically-shaped tensors to bound code growth.
- }];
-
- let dependentDialects = [
- "mlir::arith::ArithDialect",
- "mlir::linalg::LinalgDialect",
- "mlir::tensor::TensorDialect"
- ];
-}
-
def EmptyPass : Pass<"empty"> {
let summary = "Empty pass that does nothing.";
diff --git a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
index 7e4a7b40b3..01bbf4349e 100644
--- a/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
+++ b/mlir/include/Driver/DefaultPipelines/DefaultPipelines.h
@@ -71,29 +71,16 @@ const PipelineList pipelineList{
"scatter-lowering",
"hlo-custom-call-lowering",
"cse",
- // Sink scalar extractions through small-tensor producers and fuse the
- // remaining elementwise ops. Traced gate-parameter dataflow (e.g. runtime
- // Hamiltonian coefficients) otherwise reaches bufferization as thousands
- // of tiny tensor ops that each become an alloc + copy.
- "scalarize-tensor-extracts",
- "func.func(linalg-fuse-elementwise-ops)",
- "canonicalize",
// Reconstruct the loops that tracing unrolled (e.g. Trotter steps); a
// repeat of multiplicity k shrinks its region k-fold before the
- // bufferization and LLVM stages amplify it. No cse is needed before
- // this pass computations duplicated per iteration sit inside each
- // repeated window and reroll as part of it, and constants (the only
- // cross-window operands that must be identical SSA values) are already
- // uniqued by the preceding canonicalize. An ablation with cse here
- // showed no change in rerolling or output size.
+ // bufferization and LLVM stages amplify it. Computations duplicated per
+ // iteration sit inside each repeated window and reroll as part of it,
+ // and constants (the only cross-window operands that must be identical
+ // SSA values) are already uniqued by the preceding cse.
"reroll-loops",
"func.func(linalg-detensorize{aggressive-mode})",
"detensorize-scf",
"detensorize-function-boundary",
- // Detensorization is what materializes tensor.extract on the gate-angle
- // dataflow, so scalarization must run again here to fold the
- // extract_slice/collapse_shape chains it exposes.
- "scalarize-tensor-extracts",
"canonicalize",
"symbol-dce"}},
{"gradient-lowering-stage",
diff --git a/mlir/lib/Catalyst/Transforms/CMakeLists.txt b/mlir/lib/Catalyst/Transforms/CMakeLists.txt
index 2ecd8c019a..15cefe6437 100644
--- a/mlir/lib/Catalyst/Transforms/CMakeLists.txt
+++ b/mlir/lib/Catalyst/Transforms/CMakeLists.txt
@@ -24,7 +24,6 @@ file(GLOB SRC
RegisterInactiveCallbackPass.cpp
RerollLoopsPass.cpp
ResourceAnalysisPass.cpp
- ScalarizeTensorExtractsPass.cpp
RegisterDecompRuleResourcePass.cpp
SplitMultipleTapes.cpp
TBAAPatterns.cpp
diff --git a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp b/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
deleted file mode 100644
index e60aca1a96..0000000000
--- a/mlir/lib/Catalyst/Transforms/ScalarizeTensorExtractsPass.cpp
+++ /dev/null
@@ -1,292 +0,0 @@
-// Copyright 2026 Xanadu Quantum Technologies Inc.
-
-// Licensed under the Apache License, Version 2.0 (the "License");
-// you may not use this file except in compliance with the License.
-// You may obtain a copy of the License at
-
-// http://www.apache.org/licenses/LICENSE-2.0
-
-// Unless required by applicable law or agreed to in writing, software
-// distributed under the License is distributed on an "AS IS" BASIS,
-// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-// See the License for the specific language governing permissions and
-// limitations under the License.
-
-#define DEBUG_TYPE "scalarize-tensor-extracts"
-
-#include "mlir/Dialect/Arith/IR/Arith.h"
-#include "mlir/Dialect/Linalg/IR/Linalg.h"
-#include "mlir/Dialect/Tensor/IR/Tensor.h"
-#include "mlir/IR/IRMapping.h"
-#include "mlir/Interfaces/SideEffectInterfaces.h"
-#include "mlir/Pass/Pass.h"
-#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
-
-using namespace llvm;
-using namespace mlir;
-
-// Gate parameters computed from runtime inputs (e.g. Trotterization with
-// runtime Hamiltonian coefficients) arrive as long chains of small tensor ops
-// consumed only by scalar `tensor.extract` operations; each such tensor
-// survives bufferization as an allocation plus copies. This pass sinks
-// `tensor.extract` through `linalg.generic` (inlining the scalar payload),
-// `tensor.collapse_shape`, and `tensor.extract_slice`, so the extracted
-// element is computed in scalar arithmetic and the tensors become dead.
-// Payload inlining is limited to small statically-shaped results to bound
-// code growth.
-
-namespace {
-
-/// Upper bound on the number of elements of a tensor whose producer payload we
-/// are willing to clone per extraction site. 16 covers the 4x4 matrices of
-/// two-qubit gates, the largest tensors in gate-parameter dataflow, while
-/// keeping the worst-case code growth per extract small.
-constexpr int64_t kMaxScalarizedElements = 16;
-
-/// Fold tensor.extract(linalg.generic) by inlining the generic's scalar payload
-/// at the extraction point, for elementwise (all-parallel) generics.
-struct ExtractOfGeneric : public OpRewritePattern {
- using OpRewritePattern::OpRewritePattern;
-
- LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
- PatternRewriter &rewriter) const override
- {
- auto genericOp = extractOp.getTensor().getDefiningOp();
- if (!genericOp) {
- return failure();
- }
-
- // Only elementwise generics: every iterator is parallel.
- if (genericOp.getNumParallelLoops() != genericOp.getNumLoops()) {
- return failure();
- }
-
- // Restrict to small statically shaped results to bound code growth.
- auto resultType = dyn_cast(extractOp.getTensor().getType());
- if (!resultType || !resultType.hasStaticShape() ||
- resultType.getNumElements() > kMaxScalarizedElements) {
- return failure();
- }
-
- // Identify which result of the generic is being extracted and require an
- // identity indexing map for it, so the iteration indices equal the
- // extraction indices.
- auto resultNumber = cast(extractOp.getTensor()).getResultNumber();
- OpOperand *initOperand = genericOp.getDpsInitOperand(resultNumber);
- AffineMap outputMap = genericOp.getMatchingIndexingMap(initOperand);
- if (!outputMap.isIdentity()) {
- return failure();
- }
-
- Block *body = genericOp.getBody();
-
- // The payload must be speculatable scalar code and must not read the
- // accumulator (output block argument).
- for (Operation &op : body->without_terminator()) {
- if (!isPure(&op)) {
- return failure();
- }
- }
- for (OpOperand &outOperand : genericOp.getDpsInitsMutable()) {
- BlockArgument outArg = body->getArgument(outOperand.getOperandNumber());
- if (!outArg.use_empty()) {
- return failure();
- }
- }
-
- // Validate every input indexing map up front: patterns must not mutate
- // the IR on the failure path, and ops are created per input below.
- for (OpOperand *inOperand : genericOp.getDpsInputOperands()) {
- if (!isa(inOperand->get().getType())) {
- continue;
- }
- AffineMap inputMap = genericOp.getMatchingIndexingMap(inOperand);
- for (AffineExpr expr : inputMap.getResults()) {
- if (!isa(expr)) {
- return failure();
- }
- }
- }
-
- Location loc = extractOp.getLoc();
- SmallVector iterIndices(extractOp.getIndices());
-
- // Materialize scalar operands: one tensor.extract per generic input, at
- // indices given by composing that input's indexing map with the
- // extraction indices.
- IRMapping mapping;
- for (OpOperand *inOperand : genericOp.getDpsInputOperands()) {
- BlockArgument blockArg = body->getArgument(inOperand->getOperandNumber());
- Value input = inOperand->get();
- if (!isa(input.getType())) {
- // Scalar operands of the generic map through unchanged.
- mapping.map(blockArg, input);
- continue;
- }
- AffineMap inputMap = genericOp.getMatchingIndexingMap(inOperand);
- SmallVector inputIndices;
- for (AffineExpr expr : inputMap.getResults()) {
- if (auto dimExpr = dyn_cast(expr)) {
- inputIndices.push_back(iterIndices[dimExpr.getPosition()]);
- }
- else {
- auto constExpr = cast(expr);
- inputIndices.push_back(
- arith::ConstantIndexOp::create(rewriter, loc, constExpr.getValue()));
- }
- }
- Value scalar = tensor::ExtractOp::create(rewriter, loc, input, inputIndices);
- mapping.map(blockArg, scalar);
- }
-
- // Clone the payload, resolving linalg.index to the extraction indices.
- for (Operation &op : body->without_terminator()) {
- if (auto indexOp = dyn_cast(op)) {
- mapping.map(indexOp.getResult(), iterIndices[indexOp.getDim()]);
- continue;
- }
- rewriter.clone(op, mapping);
- }
-
- auto yieldOp = cast(body->getTerminator());
- Value result = mapping.lookupOrDefault(yieldOp.getOperand(resultNumber));
- rewriter.replaceOp(extractOp, result);
- return success();
- }
-};
-
-/// Fold tensor.extract(tensor.collapse_shape) for collapses that only drop or
-/// merge unit dimensions (at most one non-unit dimension per reassociation
-/// group), by extracting directly from the source.
-struct ExtractOfCollapseShape : public OpRewritePattern {
- using OpRewritePattern::OpRewritePattern;
-
- LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
- PatternRewriter &rewriter) const override
- {
- auto collapseOp = extractOp.getTensor().getDefiningOp();
- if (!collapseOp) {
- return failure();
- }
- auto srcType = collapseOp.getSrcType();
- if (!srcType.hasStaticShape()) {
- return failure();
- }
-
- // Validate every reassociation group up front (at most one non-unit
- // source dimension per group): patterns must not mutate the IR on the
- // failure path, and index constants are created per group below.
- SmallVector groups = collapseOp.getReassociationIndices();
- SmallVector nonUnitDims(groups.size(), -1);
- for (const auto &[groupIdx, group] : llvm::enumerate(groups)) {
- for (int64_t srcDim : group) {
- if (srcType.getDimSize(srcDim) != 1) {
- if (nonUnitDims[groupIdx] != -1) {
- return failure(); // true merge of two non-unit dims
- }
- nonUnitDims[groupIdx] = srcDim;
- }
- }
- }
-
- Location loc = extractOp.getLoc();
- SmallVector srcIndices(srcType.getRank());
-
- Value zero;
- auto getZero = [&]() {
- if (!zero) {
- zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
- }
- return zero;
- };
-
- // A rank-0 result means every source dimension is a unit dimension.
- if (groups.empty()) {
- for (int64_t dim = 0; dim < srcType.getRank(); ++dim) {
- srcIndices[dim] = getZero();
- }
- }
- for (const auto &[groupIdx, group] : llvm::enumerate(groups)) {
- for (int64_t srcDim : group) {
- srcIndices[srcDim] = (srcDim == nonUnitDims[groupIdx])
- ? extractOp.getIndices()[groupIdx]
- : getZero();
- }
- }
-
- rewriter.replaceOpWithNewOp(extractOp, collapseOp.getSrc(), srcIndices);
- return success();
- }
-};
-
-/// Fold tensor.extract(tensor.extract_slice) by extracting from the source at
-/// offset + index * stride.
-struct ExtractOfExtractSlice : public OpRewritePattern {
- using OpRewritePattern::OpRewritePattern;
-
- LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
- PatternRewriter &rewriter) const override
- {
- auto sliceOp = extractOp.getTensor().getDefiningOp();
- if (!sliceOp) {
- return failure();
- }
-
- Location loc = extractOp.getLoc();
- int64_t srcRank = sliceOp.getSourceType().getRank();
-
- // The slice may be rank-reducing: map each source dim to its position in
- // the result (or none if the dim was dropped).
- llvm::SmallBitVector droppedDims = sliceOp.getDroppedDims();
-
- auto materialize = [&](OpFoldResult ofr) -> Value {
- if (auto val = dyn_cast(ofr)) {
- return val;
- }
- return arith::ConstantIndexOp::create(rewriter, loc,
- cast(cast(ofr)).getInt());
- };
-
- SmallVector srcIndices;
- unsigned resultDim = 0;
- for (int64_t dim = 0; dim < srcRank; ++dim) {
- Value offset = materialize(sliceOp.getMixedOffsets()[dim]);
- if (droppedDims.test(dim)) {
- srcIndices.push_back(offset);
- continue;
- }
- Value index = extractOp.getIndices()[resultDim++];
- Value stride = materialize(sliceOp.getMixedStrides()[dim]);
- Value scaled = arith::MulIOp::create(rewriter, loc, index, stride);
- srcIndices.push_back(arith::AddIOp::create(rewriter, loc, offset, scaled));
- }
-
- rewriter.replaceOpWithNewOp(extractOp, sliceOp.getSource(), srcIndices);
- return success();
- }
-};
-
-} // namespace
-
-namespace catalyst {
-
-#define GEN_PASS_DEF_SCALARIZETENSOREXTRACTSPASS
-#include "Catalyst/Transforms/Passes.h.inc"
-
-struct ScalarizeTensorExtractsPass
- : public impl::ScalarizeTensorExtractsPassBase {
- using impl::ScalarizeTensorExtractsPassBase<
- ScalarizeTensorExtractsPass>::ScalarizeTensorExtractsPassBase;
-
- void runOnOperation() override
- {
- MLIRContext *context = &getContext();
- RewritePatternSet patterns(context);
- patterns.add(context);
- if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) {
- signalPassFailure();
- }
- }
-};
-
-} // namespace catalyst
diff --git a/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir b/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
deleted file mode 100644
index d41df225da..0000000000
--- a/mlir/test/Catalyst/ScalarizeTensorExtractsTest.mlir
+++ /dev/null
@@ -1,146 +0,0 @@
-// Copyright 2026 Xanadu Quantum Technologies Inc.
-
-// Licensed under the Apache License, Version 2.0 (the "License");
-// you may not use this file except in compliance with the License.
-// You may obtain a copy of the License at
-
-// http://www.apache.org/licenses/LICENSE-2.0
-
-// Unless required by applicable law or agreed to in writing, software
-// distributed under the License is distributed on an "AS IS" BASIS,
-// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-// See the License for the specific language governing permissions and
-// limitations under the License.
-
-// RUN: quantum-opt --scalarize-tensor-extracts --canonicalize --split-input-file %s | FileCheck %s
-
-// The extract_slice + collapse_shape + extract chain produced when tracing
-// indexes a 1-D runtime tensor collapses to a single extract.
-
-// CHECK-LABEL: @slice_collapse_extract
-// CHECK-SAME: (%[[ARG:.+]]: tensor<15xf64>)
-// CHECK: %[[C14:.+]] = arith.constant 14 : index
-// CHECK: %[[RES:.+]] = tensor.extract %[[ARG]][%[[C14]]] : tensor<15xf64>
-// CHECK-NOT: tensor.extract_slice
-// CHECK-NOT: tensor.collapse_shape
-// CHECK: return %[[RES]]
-func.func @slice_collapse_extract(%arg0: tensor<15xf64>) -> f64 {
- %s = tensor.extract_slice %arg0[14] [1] [1] : tensor<15xf64> to tensor<1xf64>
- %c = tensor.collapse_shape %s [] : tensor<1xf64> into tensor
- %e = tensor.extract %c[] : tensor
- return %e : f64
-}
-
-// -----
-
-// Extracting one element of an elementwise linalg.generic inlines the scalar
-// payload; the generic and its tensor.empty become dead.
-
-// CHECK-LABEL: @extract_of_generic
-// CHECK-SAME: (%[[A:.+]]: tensor<2x2xf64>, %[[B:.+]]: tensor<2x2xf64>)
-// CHECK-DAG: %[[C0:.+]] = arith.constant 0 : index
-// CHECK-DAG: %[[C1:.+]] = arith.constant 1 : index
-// CHECK-DAG: %[[EA:.+]] = tensor.extract %[[A]][%[[C0]], %[[C1]]]
-// CHECK-DAG: %[[EB:.+]] = tensor.extract %[[B]][%[[C0]], %[[C1]]]
-// CHECK: %[[RES:.+]] = arith.mulf %[[EA]], %[[EB]] : f64
-// CHECK-NOT: linalg.generic
-// CHECK: return %[[RES]]
-#map = affine_map<(d0, d1) -> (d0, d1)>
-func.func @extract_of_generic(%arg0: tensor<2x2xf64>, %arg1: tensor<2x2xf64>) -> f64 {
- %c0 = arith.constant 0 : index
- %c1 = arith.constant 1 : index
- %empty = tensor.empty() : tensor<2x2xf64>
- %prod = linalg.generic {indexing_maps = [#map, #map, #map], iterator_types = ["parallel", "parallel"]}
- ins(%arg0, %arg1 : tensor<2x2xf64>, tensor<2x2xf64>) outs(%empty : tensor<2x2xf64>) {
- ^bb0(%in0: f64, %in1: f64, %out: f64):
- %m = arith.mulf %in0, %in1 : f64
- linalg.yield %m : f64
- } -> tensor<2x2xf64>
- %res = tensor.extract %prod[%c0, %c1] : tensor<2x2xf64>
- return %res : f64
-}
-
-// -----
-
-// Broadcast (rank-0 to 2x2) generics fold to an extract of the rank-0 source.
-
-// CHECK-LABEL: @extract_of_broadcast
-// CHECK-SAME: (%[[A:.+]]: tensor)
-// CHECK: %[[E:.+]] = tensor.extract %[[A]][] : tensor
-// CHECK-NOT: linalg.generic
-// CHECK: return %[[E]]
-#map0 = affine_map<(d0, d1) -> ()>
-#map1 = affine_map<(d0, d1) -> (d0, d1)>
-func.func @extract_of_broadcast(%arg0: tensor) -> f64 {
- %c1 = arith.constant 1 : index
- %empty = tensor.empty() : tensor<2x2xf64>
- %bcast = linalg.generic {indexing_maps = [#map0, #map1], iterator_types = ["parallel", "parallel"]}
- ins(%arg0 : tensor) outs(%empty : tensor<2x2xf64>) {
- ^bb0(%in: f64, %out: f64):
- linalg.yield %in : f64
- } -> tensor<2x2xf64>
- %res = tensor.extract %bcast[%c1, %c1] : tensor<2x2xf64>
- return %res : f64
-}
-
-// -----
-
-// Reductions must not be scalarized: the payload reads the accumulator.
-
-// CHECK-LABEL: @reduction_untouched
-// CHECK: linalg.generic
-// CHECK: tensor.extract
-#map_in = affine_map<(d0) -> (d0)>
-#map_out = affine_map<(d0) -> ()>
-func.func @reduction_untouched(%arg0: tensor<8xf64>) -> f64 {
- %cst = arith.constant 0.0 : f64
- %empty = tensor.empty() : tensor
- %fill = linalg.fill ins(%cst : f64) outs(%empty : tensor) -> tensor
- %sum = linalg.generic {indexing_maps = [#map_in, #map_out], iterator_types = ["reduction"]}
- ins(%arg0 : tensor<8xf64>) outs(%fill : tensor) {
- ^bb0(%in: f64, %acc: f64):
- %a = arith.addf %in, %acc : f64
- linalg.yield %a : f64
- } -> tensor
- %res = tensor.extract %sum[] : tensor
- return %res : f64
-}
-
-// -----
-
-// Large tensors must not be scalarized (payload cloning is capped).
-
-// CHECK-LABEL: @large_tensor_untouched
-// CHECK: linalg.generic
-// CHECK: tensor.extract
-#map2 = affine_map<(d0) -> (d0)>
-func.func @large_tensor_untouched(%arg0: tensor<100xf64>) -> f64 {
- %c5 = arith.constant 5 : index
- %empty = tensor.empty() : tensor<100xf64>
- %sq = linalg.generic {indexing_maps = [#map2, #map2], iterator_types = ["parallel"]}
- ins(%arg0 : tensor<100xf64>) outs(%empty : tensor<100xf64>) {
- ^bb0(%in: f64, %out: f64):
- %m = arith.mulf %in, %in : f64
- linalg.yield %m : f64
- } -> tensor<100xf64>
- %res = tensor.extract %sq[%c5] : tensor<100xf64>
- return %res : f64
-}
-
-// -----
-
-// Rank-reducing extract_slice: index arithmetic offset + i * stride.
-
-// CHECK-LABEL: @strided_slice_extract
-// CHECK-SAME: (%[[ARG:.+]]: tensor<4x6xf64>, %[[I:.+]]: index)
-// CHECK-DAG: %[[C2:.+]] = arith.constant 2 : index
-// CHECK-DAG: %[[C1:.+]] = arith.constant 1 : index
-// CHECK-DAG: %[[SCALED:.+]] = arith.muli %[[I]], %[[C2]]
-// CHECK-DAG: %[[COL:.+]] = arith.addi %[[SCALED]], %[[C1]]
-// CHECK: %[[RES:.+]] = tensor.extract %[[ARG]][%[[C2]], %[[COL]]] : tensor<4x6xf64>
-// CHECK: return %[[RES]]
-func.func @strided_slice_extract(%arg0: tensor<4x6xf64>, %i: index) -> f64 {
- %s = tensor.extract_slice %arg0[2, 1] [1, 3] [1, 2] : tensor<4x6xf64> to tensor<3xf64>
- %e = tensor.extract %s[%i] : tensor<3xf64>
- return %e : f64
-}
From 5fa0e5e6db9d4987f25d1141f28d519d574e2731 Mon Sep 17 00:00:00 2001
From: jk20342 <155792753+jk20342@users.noreply.github.com>
Date: Thu, 16 Jul 2026 15:44:48 -0400
Subject: [PATCH 6/9] Update pull request link in changelog-dev.md
Updated pull request reference in changelog.
---
doc/releases/changelog-dev.md | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index 9d513614ef..2c6de282d3 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -172,7 +172,7 @@
k-fold before the bufferization and LLVM stages amplify it, reducing compile time,
peak memory, and IR size on workloads with many structurally identical circuit
segments. The pass runs by default in the HLO lowering stage.
- [(#XXXX)](https://github.com/PennyLaneAI/catalyst/pull/XXXX)
+ [(#3036)](https://github.com/PennyLaneAI/catalyst/pull/3036)
Breaking changes 💔
From c04737b5b41dc584377c86bbc58b185dfd90a7a0 Mon Sep 17 00:00:00 2001
From: jk20342 <155792753+jk20342@users.noreply.github.com>
Date: Thu, 16 Jul 2026 15:48:58 -0400
Subject: [PATCH 7/9] Update changelog-dev.md
---
doc/releases/changelog-dev.md | 14 ++++++++------
1 file changed, 8 insertions(+), 6 deletions(-)
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index 2c6de282d3..69eeb128f6 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -166,12 +166,14 @@
* Added ``CZ`` support to ``to-ppr`` pass.
[(#3009)](https://github.com/PennyLaneAI/catalyst/pull/3009)
-* A new `reroll-loops` pass reconstructs the loops that tracing unrolled by detecting
- repeated op sequences (e.g. Trotter steps, layers, folds) via structural hashing and
- rewriting them as `scf.for` loops. A repeat of multiplicity k shrinks that IR region
- k-fold before the bufferization and LLVM stages amplify it, reducing compile time,
- peak memory, and IR size on workloads with many structurally identical circuit
- segments. The pass runs by default in the HLO lowering stage.
+* IR amplification for circuits whose gate parameters are computed from runtime values
+ (e.g. `qml.TrotterProduct` with runtime Hamiltonian coefficients) is drastically reduced.
+ The new `scalarize-tensor-extracts` pass turns gate-angle dataflow into scalar arithmetic
+ instead of thousands of small tensors that each survive bufferization as an allocation,
+ and the new `reroll-loops` pass reconstructs the loops that tracing unrolled by rewriting
+ repeated op sequences as `scf.for` loops. On a Trotterized QPE workload with runtime
+ coefficients, compile time, peak memory, and final IR size all drop by large factors.
+ [(#3013)](https://github.com/PennyLaneAI/catalyst/pull/3013)
[(#3036)](https://github.com/PennyLaneAI/catalyst/pull/3036)
Breaking changes 💔
From 6685ceacdc972bbdab1d32a709b6888e54a2cfe3 Mon Sep 17 00:00:00 2001
From: jk20342 <155792753+jk20342@users.noreply.github.com>
Date: Fri, 17 Jul 2026 03:57:07 -0400
Subject: [PATCH 8/9] Update changelog-dev.md
---
doc/releases/changelog-dev.md | 1 -
1 file changed, 1 deletion(-)
diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md
index 69eeb128f6..2dd337be02 100644
--- a/doc/releases/changelog-dev.md
+++ b/doc/releases/changelog-dev.md
@@ -173,7 +173,6 @@
and the new `reroll-loops` pass reconstructs the loops that tracing unrolled by rewriting
repeated op sequences as `scf.for` loops. On a Trotterized QPE workload with runtime
coefficients, compile time, peak memory, and final IR size all drop by large factors.
- [(#3013)](https://github.com/PennyLaneAI/catalyst/pull/3013)
[(#3036)](https://github.com/PennyLaneAI/catalyst/pull/3036)
Breaking changes 💔
From 38a4d9624ac9099661ea243554e13586c0dd545e Mon Sep 17 00:00:00 2001
From: Jacob Kitchen <155792753+JakeKitchen@users.noreply.github.com>
Date: Tue, 21 Jul 2026 16:12:32 -0400
Subject: [PATCH 9/9] resolve codefactor
---
frontend/test/pytest/test_reroll_loops.py | 63 +++++++++++++++--------
1 file changed, 41 insertions(+), 22 deletions(-)
diff --git a/frontend/test/pytest/test_reroll_loops.py b/frontend/test/pytest/test_reroll_loops.py
index f1c35672d7..852190572b 100644
--- a/frontend/test/pytest/test_reroll_loops.py
+++ b/frontend/test/pytest/test_reroll_loops.py
@@ -24,26 +24,50 @@
from jax import numpy as jnp
from catalyst import qjit
+from catalyst.debug import get_compilation_stage
# H2/STO-3G-like coefficients; the structure (15 terms, runtime values) is what
# exercises the pipeline, the values just need to be a valid Hamiltonian.
COEFFS = [
- -0.0996, 0.1711, 0.1711, -0.2225, -0.2225, 0.1686, 0.0453, -0.0453,
- -0.0453, 0.0453, 0.1205, 0.1658, 0.1658, 0.1205, 0.1743,
-]
-OPS_FACTORY = lambda: [
- qml.Identity(0),
- qml.PauliZ(0), qml.PauliZ(1), qml.PauliZ(2), qml.PauliZ(3),
- qml.PauliZ(0) @ qml.PauliZ(1),
- qml.PauliY(0) @ qml.PauliX(1) @ qml.PauliX(2) @ qml.PauliY(3),
- qml.PauliY(0) @ qml.PauliY(1) @ qml.PauliX(2) @ qml.PauliX(3),
- qml.PauliX(0) @ qml.PauliX(1) @ qml.PauliY(2) @ qml.PauliY(3),
- qml.PauliX(0) @ qml.PauliY(1) @ qml.PauliY(2) @ qml.PauliX(3),
- qml.PauliZ(0) @ qml.PauliZ(2), qml.PauliZ(0) @ qml.PauliZ(3),
- qml.PauliZ(1) @ qml.PauliZ(2), qml.PauliZ(1) @ qml.PauliZ(3),
- qml.PauliZ(2) @ qml.PauliZ(3),
+ -0.0996,
+ 0.1711,
+ 0.1711,
+ -0.2225,
+ -0.2225,
+ 0.1686,
+ 0.0453,
+ -0.0453,
+ -0.0453,
+ 0.0453,
+ 0.1205,
+ 0.1658,
+ 0.1658,
+ 0.1205,
+ 0.1743,
]
+
+def make_ops():
+ """Build a fresh list of Hamiltonian terms (operators can't be reused)."""
+ return [
+ qml.Identity(0),
+ qml.PauliZ(0),
+ qml.PauliZ(1),
+ qml.PauliZ(2),
+ qml.PauliZ(3),
+ qml.PauliZ(0) @ qml.PauliZ(1),
+ qml.PauliY(0) @ qml.PauliX(1) @ qml.PauliX(2) @ qml.PauliY(3),
+ qml.PauliY(0) @ qml.PauliY(1) @ qml.PauliX(2) @ qml.PauliX(3),
+ qml.PauliX(0) @ qml.PauliX(1) @ qml.PauliY(2) @ qml.PauliY(3),
+ qml.PauliX(0) @ qml.PauliY(1) @ qml.PauliY(2) @ qml.PauliX(3),
+ qml.PauliZ(0) @ qml.PauliZ(2),
+ qml.PauliZ(0) @ qml.PauliZ(3),
+ qml.PauliZ(1) @ qml.PauliZ(2),
+ qml.PauliZ(1) @ qml.PauliZ(3),
+ qml.PauliZ(2) @ qml.PauliZ(3),
+ ]
+
+
N_QUBITS = 4
N_EST = 2
N_TROTTER = 6 # enough repetitions for reroll-loops to fire
@@ -58,13 +82,12 @@ def qpe_circuit(coeffs):
qml.PauliX(1)
for k in range(N_EST):
qml.Hadamard(wires=N_QUBITS + k)
- H = qml.dot(coeffs if runtime else COEFFS, OPS_FACTORY())
+ H = qml.dot(coeffs if runtime else COEFFS, make_ops())
for k in range(N_EST):
t = 2 ** (N_EST - 1 - k)
qml.ctrl(
qml.adjoint(
- qml.TrotterProduct(H, time=t, n=N_TROTTER, order=2,
- check_hermitian=False)
+ qml.TrotterProduct(H, time=t, n=N_TROTTER, order=2, check_hermitian=False)
),
control=N_QUBITS + k,
)
@@ -94,8 +117,6 @@ def test_reroll_recovers_loops(self):
the unrolled Trotter steps, and fewer gate ops than the unrolled
circuit (guards against silent regression of reroll-loops in the
default pipeline)."""
- from catalyst.debug import get_compilation_stage
-
dev = qml.device("lightning.qubit", wires=N_QUBITS + N_EST)
coeffs = jnp.array(COEFFS)
@@ -106,9 +127,7 @@ def test_reroll_recovers_loops(self):
lowered = get_compilation_stage(compiled, "HLOLoweringStage")
# Rerolled Trotter steps are scf.for loops threading qubit values
# through iter_args.
- qubit_loops = re.findall(
- r"scf\.for .*iter_args\([^)]*\).*->.*!quantum\.bit", lowered
- )
+ qubit_loops = re.findall(r"scf\.for .*iter_args\([^)]*\).*->.*!quantum\.bit", lowered)
assert qubit_loops, "reroll-loops did not produce qubit-threading loops"
unrolled_gates = traced.count("quantum.custom")
rerolled_gates = lowered.count("quantum.custom")