From 3b31d8cacac7b632c8f8c69f2978394ef0e25b3a Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 17:50:33 +0200 Subject: [PATCH 1/3] =?UTF-8?q?=F0=9F=90=9B=20Make=20QTensor=20shrinking?= =?UTF-8?q?=20sparse=20and=20atomic=20(#2255)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 Sol via Codex --- .../QTensor/Transforms/ShrinkRegisters.cpp | 186 ++++++------------ .../Transforms/test_qtensor_transforms.cpp | 79 ++++++++ 2 files changed, 142 insertions(+), 123 deletions(-) diff --git a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp index cba3aba988..88acb015e1 100644 --- a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp +++ b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp @@ -11,6 +11,8 @@ #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include "mlir/Dialect/QTensor/Transforms/Passes.h" +#include +#include #include #include #include @@ -21,7 +23,6 @@ #include #include -#include #include #include #include @@ -31,63 +32,36 @@ namespace mlir::qtensor { #define GEN_PASS_DEF_SHRINKQTENSORTOFITPASS #include "mlir/Dialect/QTensor/Transforms/Passes.h.inc" -/** - * @brief Return the unique user of a linear qtensor value. - */ -[[nodiscard]] static Operation* getLinearTensorUser(Value tensor) { - assert(tensor.hasOneUse() && "Expected a linear tensor with exactly one use"); - return *tensor.getUsers().begin(); -} - /** * @brief Mark a single live index. */ -[[nodiscard]] static LogicalResult markLiveIndex(const int64_t index, - BitVector& liveIndices) { - if (index < 0 || std::cmp_greater_equal(index, liveIndices.size())) { +[[nodiscard]] static LogicalResult +markLiveIndex(int64_t index, int64_t tensorSize, + llvm::SmallDenseSet& liveIndices) { + if (index < 0 || index >= tensorSize) { return failure(); } - liveIndices.set(static_cast(index)); + liveIndices.insert(index); return success(); } -/** - * @brief Redirect the tensor operand from @p from to @p to. - */ -[[nodiscard]] static LogicalResult remapTensorOperand(Operation* op, Value from, - Value to) { - if (auto extractOp = dyn_cast(op)) { - if (extractOp.getTensor() != from) { - return failure(); - } - extractOp->setOperand(0, to); - return success(); - } - if (auto insertOp = dyn_cast(op)) { - if (insertOp.getDest() != from) { - return failure(); - } - insertOp->setOperand(1, to); - return success(); - } - if (auto deallocOp = dyn_cast(op)) { - if (deallocOp.getTensor() != from) { - return failure(); - } - deallocOp->setOperand(0, to); - return success(); - } - return failure(); -} +struct TensorAccess { + Operation* operation; + int64_t index; +}; /** - * @brief Walk alloc->dealloc and collect all touched indices. + * @brief Walk alloc->dealloc and plan all accesses without changing the IR. */ -[[nodiscard]] static LogicalResult -collectLiveIndices(AllocOp allocOp, BitVector& live, DeallocOp& deallocOp) { +[[nodiscard]] static LogicalResult collectTensorChain( + AllocOp allocOp, int64_t tensorSize, llvm::SmallDenseSet& live, + SmallVectorImpl& accesses, DeallocOp& deallocOp) { auto tensor = allocOp.getResult(); while (true) { - auto* user = getLinearTensorUser(tensor); + if (!tensor.hasOneUse()) { + return failure(); + } + auto* user = *tensor.getUsers().begin(); if (auto currentDealloc = dyn_cast(user)) { if (currentDealloc.getTensor() != tensor) { @@ -102,9 +76,10 @@ collectLiveIndices(AllocOp allocOp, BitVector& live, DeallocOp& deallocOp) { return failure(); } auto index = getConstantIntValue(extractOp.getIndex()); - if (!index || failed(markLiveIndex(*index, live))) { + if (!index || failed(markLiveIndex(*index, tensorSize, live))) { return failure(); } + accesses.push_back({extractOp, *index}); tensor = extractOp.getOutTensor(); continue; } @@ -114,9 +89,10 @@ collectLiveIndices(AllocOp allocOp, BitVector& live, DeallocOp& deallocOp) { return failure(); } auto index = getConstantIntValue(insertOp.getIndex()); - if (!index || failed(markLiveIndex(*index, live))) { + if (!index || failed(markLiveIndex(*index, tensorSize, live))) { return failure(); } + accesses.push_back({insertOp, *index}); tensor = insertOp.getResult(); continue; } @@ -141,9 +117,11 @@ struct ShrinkStaticQTensor final : OpRewritePattern { return failure(); } - BitVector live(static_cast(*oldSize), false); + llvm::SmallDenseSet live; + SmallVector accesses; DeallocOp oldDeallocOp{}; - if (failed(collectLiveIndices(allocOp, live, oldDeallocOp))) { + if (failed(collectTensorChain(allocOp, *oldSize, live, accesses, + oldDeallocOp))) { return failure(); } @@ -151,57 +129,41 @@ struct ShrinkStaticQTensor final : OpRewritePattern { return failure(); } - SmallVector newIndexByOldIndex(static_cast(*oldSize), -1); - int64_t newSize = 0; - for (int64_t index = 0; index < *oldSize; ++index) { - if (live.test(static_cast(index))) { - newIndexByOldIndex[static_cast(index)] = newSize++; - } + SmallVector liveIndices(live.begin(), live.end()); + llvm::sort(liveIndices); + const auto newSize = static_cast(liveIndices.size()); + DenseMap newIndexByOldIndex; + for (auto [newIndex, oldIndex] : llvm::enumerate(liveIndices)) { + newIndexByOldIndex.try_emplace(oldIndex, static_cast(newIndex)); } if (newSize <= 0 || newSize == *oldSize) { return failure(); } + SmallVector mappedIndices; + mappedIndices.reserve(accesses.size()); + for (const auto& access : accesses) { + const auto mapped = newIndexByOldIndex.find(access.index); + if (mapped == newIndexByOldIndex.end()) { + return failure(); + } + mappedIndices.push_back(mapped->second); + } + rewriter.setInsertionPoint(allocOp); auto size = arith::ConstantIndexOp::create(rewriter, allocOp.getLoc(), newSize); auto newAlloc = AllocOp::create(rewriter, allocOp.getLoc(), size.getResult()); - newAlloc->setDiscardableAttrs(allocOp->getDiscardableAttrDictionary()); + rewriter.modifyOpInPlace(newAlloc, [&] { + newAlloc->setDiscardableAttrs(allocOp->getDiscardableAttrDictionary()); + }); - auto oldTensor = allocOp.getResult(); auto currentTensor = newAlloc.getResult(); - while (true) { - Operation* currentOp = getLinearTensorUser(oldTensor); - - if (auto deallocOp = dyn_cast(currentOp)) { - if (deallocOp != oldDeallocOp || deallocOp.getTensor() != oldTensor) { - return failure(); - } - rewriter.setInsertionPoint(deallocOp); - DeallocOp::create(rewriter, deallocOp.getLoc(), currentTensor); - rewriter.eraseOp(deallocOp); - break; - } - - if (auto extractOp = dyn_cast(currentOp)) { - if (extractOp.getTensor() != oldTensor) { - return failure(); - } - const auto oldIndex = *getConstantIntValue(extractOp.getIndex()); - if (oldIndex < 0 || - std::cmp_greater_equal(oldIndex, newIndexByOldIndex.size())) { - return failure(); - } - const auto mappedIndex = - newIndexByOldIndex[static_cast(oldIndex)]; - if (mappedIndex < 0) { - return failure(); - } - auto oldOutTensor = extractOp.getOutTensor(); - auto* nextOp = getLinearTensorUser(oldOutTensor); - + for (const auto [access, mappedIndex] : + llvm::zip_equal(accesses, mappedIndices)) { + if (auto extractOp = dyn_cast(access.operation)) { rewriter.setInsertionPoint(extractOp); auto index = arith::ConstantIndexOp::create( rewriter, extractOp.getLoc(), mappedIndex); @@ -209,50 +171,28 @@ struct ShrinkStaticQTensor final : OpRewritePattern { currentTensor, index.getResult()); rewriter.replaceAllUsesWith(extractOp.getResult(), newExtract.getResult()); - currentTensor = newExtract.getOutTensor(); - if (failed(remapTensorOperand(nextOp, oldOutTensor, oldTensor))) { - return failure(); - } - rewriter.eraseOp(extractOp); continue; } - if (auto insertOp = dyn_cast(currentOp)) { - if (insertOp.getDest() != oldTensor) { - return failure(); - } - const auto oldIndex = *getConstantIntValue(insertOp.getIndex()); - if (oldIndex < 0 || - std::cmp_greater_equal(oldIndex, newIndexByOldIndex.size())) { - return failure(); - } - const auto mappedIndex = - newIndexByOldIndex[static_cast(oldIndex)]; - if (mappedIndex < 0) { - return failure(); - } - auto oldResultTensor = insertOp.getResult(); - auto* nextOp = getLinearTensorUser(oldResultTensor); + auto insertOp = cast(access.operation); + rewriter.setInsertionPoint(insertOp); + auto index = arith::ConstantIndexOp::create(rewriter, insertOp.getLoc(), + mappedIndex); + auto newInsert = + InsertOp::create(rewriter, insertOp.getLoc(), insertOp.getScalar(), + currentTensor, index.getResult()); - rewriter.setInsertionPoint(insertOp); - auto index = arith::ConstantIndexOp::create(rewriter, insertOp.getLoc(), - mappedIndex); - auto newInsert = - InsertOp::create(rewriter, insertOp.getLoc(), insertOp.getScalar(), - currentTensor, index.getResult()); + currentTensor = newInsert.getResult(); + } - currentTensor = newInsert.getResult(); - if (failed(remapTensorOperand(nextOp, oldResultTensor, oldTensor))) { - return failure(); - } - rewriter.eraseOp(insertOp); - continue; - } + rewriter.setInsertionPoint(oldDeallocOp); + DeallocOp::create(rewriter, oldDeallocOp.getLoc(), currentTensor); - return failure(); + rewriter.eraseOp(oldDeallocOp); + for (const auto& access : llvm::reverse(accesses)) { + rewriter.eraseOp(access.operation); } - rewriter.eraseOp(allocOp); return success(); } diff --git a/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp b/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp index b16c3a8ed1..b79795f23a 100644 --- a/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp +++ b/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp @@ -20,6 +20,7 @@ #include "mlir/Dialect/QTensor/Transforms/Passes.h" #include +#include #include #include #include @@ -34,6 +35,7 @@ #include #include +#include using namespace mlir; @@ -78,4 +80,81 @@ TEST(QTensorTransformsTest, ShrinkToFitPreservesMetadata) { mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()), StringAttr::get(&context, "q")); } + +TEST(QTensorTransformsTest, HugeDeclaredTensorUsesSparseShrinkPlan) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + auto moduleOp = parseSourceString(R"mlir( + module { + func.func @main() { + %c7 = arith.constant 7 : index + %huge = arith.constant 1099511627776 : index + %reg = qtensor.alloc(%huge) : tensor<1099511627776x!qco.qubit> + %rest, %qubit = qtensor.extract %reg[%c7] + : tensor<1099511627776x!qco.qubit> + %flipped = qco.x %qubit : !qco.qubit -> !qco.qubit + %updated = qtensor.insert %flipped into %rest[%c7] + : tensor<1099511627776x!qco.qubit> + qtensor.dealloc %updated : tensor<1099511627776x!qco.qubit> + return + } + } + )mlir", + &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + PassManager manager(&context); + manager.addPass(qtensor::createShrinkQTensorToFitPass()); + ASSERT_TRUE(succeeded(manager.run(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + qtensor::AllocOp allocation; + moduleOp->walk([&](qtensor::AllocOp op) { allocation = op; }); + ASSERT_TRUE(allocation); + EXPECT_EQ(cast(allocation.getType()).getShape(), + ArrayRef{1}); +} + +TEST(QTensorTransformsTest, RejectsNonLinearChainWithoutMutation) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + auto moduleOp = parseSourceString(R"mlir( + module { + func.func @main() { + %c1 = arith.constant 1 : index + %c3 = arith.constant 3 : index + %reg = qtensor.alloc(%c3) : tensor<3x!qco.qubit> + %rest, %qubit = qtensor.extract %reg[%c1] + : tensor<3x!qco.qubit> + qtensor.dealloc %rest : tensor<3x!qco.qubit> + qtensor.dealloc %rest : tensor<3x!qco.qubit> + return + } + } + )mlir", + &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + std::string before; + llvm::raw_string_ostream(before) << *moduleOp; + + PassManager manager(&context); + manager.addPass(qtensor::createShrinkQTensorToFitPass()); + ASSERT_TRUE(succeeded(manager.run(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + std::string after; + llvm::raw_string_ostream(after) << *moduleOp; + EXPECT_EQ(after, before); +} } // namespace From 9773e0e8fdbb7bfca97141288c39814418182812 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 21:36:50 +0200 Subject: [PATCH 2/3] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Address=20QTensor=20sh?= =?UTF-8?q?rink=20review=20feedback?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 Sol via Codex --- .../QTensor/Transforms/ShrinkRegisters.cpp | 7 +--- .../Transforms/test_qtensor_transforms.cpp | 39 ------------------- 2 files changed, 2 insertions(+), 44 deletions(-) diff --git a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp index 88acb015e1..d3e54133db 100644 --- a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp +++ b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp @@ -32,6 +32,8 @@ namespace mlir::qtensor { #define GEN_PASS_DEF_SHRINKQTENSORTOFITPASS #include "mlir/Dialect/QTensor/Transforms/Passes.h.inc" +namespace { + /** * @brief Mark a single live index. */ @@ -58,9 +60,6 @@ struct TensorAccess { SmallVectorImpl& accesses, DeallocOp& deallocOp) { auto tensor = allocOp.getResult(); while (true) { - if (!tensor.hasOneUse()) { - return failure(); - } auto* user = *tensor.getUsers().begin(); if (auto currentDealloc = dyn_cast(user)) { @@ -101,8 +100,6 @@ struct TensorAccess { } } -namespace { - /** * @brief Shrink static qtensors by removing never-accessed indices. * @details QTensor is linear, so this rewrite follows a single use-def chain. diff --git a/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp b/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp index b79795f23a..b99dae7351 100644 --- a/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp +++ b/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp @@ -20,7 +20,6 @@ #include "mlir/Dialect/QTensor/Transforms/Passes.h" #include -#include #include #include #include @@ -35,7 +34,6 @@ #include #include -#include using namespace mlir; @@ -120,41 +118,4 @@ TEST(QTensorTransformsTest, HugeDeclaredTensorUsesSparseShrinkPlan) { ArrayRef{1}); } -TEST(QTensorTransformsTest, RejectsNonLinearChainWithoutMutation) { - DialectRegistry registry; - registry.insert(); - MLIRContext context(registry); - context.loadAllAvailableDialects(); - - auto moduleOp = parseSourceString(R"mlir( - module { - func.func @main() { - %c1 = arith.constant 1 : index - %c3 = arith.constant 3 : index - %reg = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - %rest, %qubit = qtensor.extract %reg[%c1] - : tensor<3x!qco.qubit> - qtensor.dealloc %rest : tensor<3x!qco.qubit> - qtensor.dealloc %rest : tensor<3x!qco.qubit> - return - } - } - )mlir", - &context); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - std::string before; - llvm::raw_string_ostream(before) << *moduleOp; - - PassManager manager(&context); - manager.addPass(qtensor::createShrinkQTensorToFitPass()); - ASSERT_TRUE(succeeded(manager.run(*moduleOp))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - std::string after; - llvm::raw_string_ostream(after) << *moduleOp; - EXPECT_EQ(after, before); -} } // namespace From 74600802521c234cf02b85330eb8db101149f45c Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 22:36:51 +0200 Subject: [PATCH 3/3] =?UTF-8?q?=F0=9F=94=A7=20Satisfy=20QTensor=20helper?= =?UTF-8?q?=20linkage=20lint?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 Sol via Codex --- mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp index d3e54133db..550d37583c 100644 --- a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp +++ b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp @@ -32,8 +32,6 @@ namespace mlir::qtensor { #define GEN_PASS_DEF_SHRINKQTENSORTOFITPASS #include "mlir/Dialect/QTensor/Transforms/Passes.h.inc" -namespace { - /** * @brief Mark a single live index. */ @@ -47,11 +45,15 @@ markLiveIndex(int64_t index, int64_t tensorSize, return success(); } +namespace { + struct TensorAccess { Operation* operation; int64_t index; }; +} // namespace + /** * @brief Walk alloc->dealloc and plan all accesses without changing the IR. */ @@ -100,6 +102,8 @@ struct TensorAccess { } } +namespace { + /** * @brief Shrink static qtensors by removing never-accessed indices. * @details QTensor is linear, so this rewrite follows a single use-def chain.