diff --git a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp index 550d37583c..0d79fc2cea 100644 --- a/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp +++ b/mlir/lib/Dialect/QTensor/Transforms/ShrinkRegisters.cpp @@ -62,6 +62,9 @@ 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)) { diff --git a/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp b/mlir/unittests/Dialect/QTensor/Transforms/test_qtensor_transforms.cpp index b99dae7351..259bbd7f15 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; @@ -118,4 +120,42 @@ 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