diff --git a/mlir/include/mlir/Dialect/QTensor/IR/QTensorDialect.td b/mlir/include/mlir/Dialect/QTensor/IR/QTensorDialect.td index 8d01d8a758..7c9864232e 100644 --- a/mlir/include/mlir/Dialect/QTensor/IR/QTensorDialect.td +++ b/mlir/include/mlir/Dialect/QTensor/IR/QTensorDialect.td @@ -23,7 +23,10 @@ def QTensorDialect : Dialect { In addition, alloc/dealloc operations are added to the dialect to support the bulk allocation and deallocation of qubit tensors with linear types. }]; - let dependentDialects = ["::mlir::qco::QCODialect"]; + let dependentDialects = ["::mlir::qco::QCODialect", + "::mlir::arith::ArithDialect"]; + + let hasCanonicalizer = 1; let cppNamespace = "::mlir::qtensor"; } diff --git a/mlir/lib/Dialect/QCO/IR/SCF/IfOp.cpp b/mlir/lib/Dialect/QCO/IR/SCF/IfOp.cpp index 8a272fefce..4358dc219c 100644 --- a/mlir/lib/Dialect/QCO/IR/SCF/IfOp.cpp +++ b/mlir/lib/Dialect/QCO/IR/SCF/IfOp.cpp @@ -10,17 +10,14 @@ #include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QCO/QCOUtils.h" -#include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include -#include #include #include #include #include #include #include -#include #include #include #include @@ -37,7 +34,6 @@ #include #include #include -#include using namespace mlir; using namespace mlir::qco; @@ -334,244 +330,12 @@ struct RemoveUnusedClassicalResults : public OpRewritePattern { } }; -struct QTensorAccess { - qtensor::ExtractOp extract; - qtensor::InsertOp insert; -}; - -struct BranchQTensorAccesses { - DenseMap accesses; - SmallVector qTensorOperations; -}; - -} // namespace - -/// Analyze a QTensor's complete lifetime in one branch. -/// -/// Supported branches extract distinct constant-index qubits, perform -/// QTensor-independent computation, reinsert one qubit at every extracted -/// index, and yield the resulting QTensor. Dynamic indices, repeated accesses, -/// and partial updates do not match. -static std::optional -analyzeQTensorBranch(Block* block, size_t qTensorArgumentIndex, - size_t qTensorYieldIndex) { - BranchQTensorAccesses result; - Value currentQTensor = block->getArgument(qTensorArgumentIndex); - bool reachedInsertPhase = false; - - while (true) { - assert(currentQTensor.hasOneUse() && "expected linear typing"); - Operation* user = *currentQTensor.getUsers().begin(); - if (user->getBlock() != block) { - return std::nullopt; - } - - if (auto extract = dyn_cast(user)) { - auto index = getConstantIntValue(extract.getIndex()); - if (reachedInsertPhase || !index || - !result.accesses - .try_emplace(*index, QTensorAccess{.extract = extract}) - .second) { - return std::nullopt; - } - result.qTensorOperations.push_back(user); - currentQTensor = extract.getOutTensor(); - continue; - } - - if (auto insert = dyn_cast(user)) { - reachedInsertPhase = true; - auto index = getConstantIntValue(insert.getIndex()); - if (!index) { - return std::nullopt; - } - auto access = result.accesses.find(*index); - if (access == result.accesses.end() || access->second.insert) { - return std::nullopt; - } - access->second.insert = insert; - result.qTensorOperations.push_back(user); - currentQTensor = insert.getResult(); - continue; - } - - auto yield = dyn_cast(user); - if (!yield || user != block->getTerminator() || - qTensorYieldIndex >= yield.getTargets().size() || - yield.getTargets()[qTensorYieldIndex] != currentQTensor || - llvm::any_of(result.accesses, [](const auto& access) { - return !access.second.insert; - })) { - return std::nullopt; - } - return result; - } -} - -/// Move a branch while replacing QTensor accesses with scalar qubits. -static void moveScalarizedQTensorBranch(IfOp oldIf, Block* oldBlock, - Block* newBlock, - size_t qTensorArgumentIndex, - BranchQTensorAccesses& accesses, - ArrayRef indices, - PatternRewriter& rewriter) { - auto oldYield = cast(oldBlock->getTerminator()); - auto scalarArguments = newBlock->getArguments().take_back(indices.size()); - auto carriedArguments = newBlock->getArguments().drop_back(indices.size()); - - SmallVector argumentReplacements; - argumentReplacements.reserve(oldBlock->getNumArguments()); - size_t carriedIndex = 0; - for (size_t oldIndex : llvm::seq(oldBlock->getNumArguments())) { - argumentReplacements.push_back(oldIndex == qTensorArgumentIndex - ? oldIf.getQubits()[qTensorArgumentIndex] - : carriedArguments[carriedIndex++]); - } - assert(carriedIndex == carriedArguments.size()); - rewriter.mergeBlocks(oldBlock, newBlock, argumentReplacements); - - SmallVector scalarYields; - scalarYields.reserve(indices.size()); - for (auto [indexPosition, index] : llvm::enumerate(indices)) { - auto access = accesses.accesses.find(index); - if (access == accesses.accesses.end()) { - scalarYields.push_back(scalarArguments[indexPosition]); - } else { - rewriter.replaceAllUsesWith(access->second.extract.getResult(), - scalarArguments[indexPosition]); - scalarYields.push_back(access->second.insert.getScalar()); - } - } - - auto oldTargets = oldYield.getTargets(); - size_t classicalResultCount = oldIf.getClassicalResults().size(); - SmallVector newYieldValues; - newYieldValues.reserve(oldTargets.size() - 1 + scalarYields.size()); - llvm::append_range(newYieldValues, - oldTargets.take_front(classicalResultCount)); - for (auto [oldIndex, value] : - llvm::enumerate(oldTargets.drop_front(classicalResultCount))) { - if (oldIndex != qTensorArgumentIndex) { - newYieldValues.push_back(value); - } - } - llvm::append_range(newYieldValues, scalarYields); - - rewriter.setInsertionPoint(oldYield); - rewriter.replaceOpWithNewOp(oldYield, newYieldValues); - - for (Operation* operation : llvm::reverse(accesses.qTensorOperations)) { - rewriter.eraseOp(operation); - } -} - -namespace { - -/// Replace constant-index QTensor updates in an if with scalar threading. -/// -/// A QTensor carried through an if hides its qubits from target mapping. This -/// pattern extracts the union of constant indices accessed by either branch, -/// threads those qubits through both branches, and reinserts the results. -/// Untouched elements remain in the QTensor outside the if. -struct ScalarizeQTensorInputs final : OpRewritePattern { - using OpRewritePattern::OpRewritePattern; - - LogicalResult matchAndRewrite(IfOp op, - PatternRewriter& rewriter) const override { - size_t classicalResultCount = op.getClassicalResults().size(); - auto oldQubits = op.getQubits(); - - for (auto [qTensorIndex, qTensor] : llvm::enumerate(oldQubits)) { - auto qTensorType = dyn_cast(qTensor.getType()); - if (!qTensorType || !qTensorType.hasStaticShape()) { - continue; - } - - auto thenAccesses = analyzeQTensorBranch( - op.thenBlock(), qTensorIndex, classicalResultCount + qTensorIndex); - auto elseAccesses = analyzeQTensorBranch( - op.elseBlock(), qTensorIndex, classicalResultCount + qTensorIndex); - if (!thenAccesses || !elseAccesses) { - continue; - } - - SmallVector accessedIndices(thenAccesses->accesses.keys()); - llvm::append_range(accessedIndices, elseAccesses->accesses.keys()); - llvm::sort(accessedIndices); - accessedIndices.erase(llvm::unique(accessedIndices), - accessedIndices.end()); - ArrayRef indices(accessedIndices); - - rewriter.setInsertionPoint(op); - SmallVector indexValues; - SmallVector scalarInputs; - indexValues.reserve(indices.size()); - scalarInputs.reserve(indices.size()); - Value qTensorWithoutScalars = qTensor; - for (int64_t index : indices) { - auto indexValue = - arith::ConstantIndexOp::create(rewriter, op.getLoc(), index); - auto extract = qtensor::ExtractOp::create(rewriter, op.getLoc(), - qTensorWithoutScalars, - indexValue.getResult()); - indexValues.push_back(indexValue.getResult()); - scalarInputs.push_back(extract.getResult()); - qTensorWithoutScalars = extract.getOutTensor(); - } - - SmallVector newQubits(oldQubits); - newQubits.erase(newQubits.begin() + qTensorIndex); - llvm::append_range(newQubits, scalarInputs); - - auto newIf = IfOp::create( - rewriter, op.getLoc(), op.getClassicalResults().getTypes(), - ValueRange(newQubits).getTypes(), op.getCondition(), newQubits); - newIf->setDiscardableAttrs(op->getDiscardableAttrDictionary()); - - SmallVector locations(newQubits.size(), op.getLoc()); - Block* oldThenBlock = op.thenBlock(); - Block* oldElseBlock = op.elseBlock(); - Block* newThenBlock = - rewriter.createBlock(&newIf.getThenRegion(), {}, - ValueRange(newQubits).getTypes(), locations); - Block* newElseBlock = - rewriter.createBlock(&newIf.getElseRegion(), {}, - ValueRange(newQubits).getTypes(), locations); - moveScalarizedQTensorBranch(op, oldThenBlock, newThenBlock, qTensorIndex, - *thenAccesses, indices, rewriter); - moveScalarizedQTensorBranch(op, oldElseBlock, newElseBlock, qTensorIndex, - *elseAccesses, indices, rewriter); - - rewriter.setInsertionPointAfter(newIf); - Value updatedQTensor = qTensorWithoutScalars; - auto scalarResults = newIf.getLinearResults().take_back(indices.size()); - for (auto [scalar, indexValue] : - llvm::zip_equal(scalarResults, indexValues)) { - updatedQTensor = - qtensor::InsertOp::create(rewriter, op.getLoc(), scalar, - updatedQTensor, indexValue) - .getResult(); - } - - SmallVector replacements( - newIf.getLinearResults().drop_back(indices.size())); - replacements.insert(replacements.begin() + qTensorIndex, updatedQTensor); - replacements.insert(replacements.begin(), - newIf.getClassicalResults().begin(), - newIf.getClassicalResults().end()); - rewriter.replaceOp(op, replacements); - return success(); - } - return failure(); - } -}; } // namespace void IfOp::getCanonicalizationPatterns(RewritePatternSet& results, MLIRContext* context) { - results - .add(context); + results.add(context); } LogicalResult IfOp::verify() { diff --git a/mlir/lib/Dialect/QIR/Execution/Runtime/CMakeLists.txt b/mlir/lib/Dialect/QIR/Execution/Runtime/CMakeLists.txt index 07935aa096..3bf6f2a7f2 100644 --- a/mlir/lib/Dialect/QIR/Execution/Runtime/CMakeLists.txt +++ b/mlir/lib/Dialect/QIR/Execution/Runtime/CMakeLists.txt @@ -21,5 +21,5 @@ if(NOT TARGET ${TARGET_NAME}) target_link_libraries( ${TARGET_NAME} PUBLIC MLIRQCODDAdapter MQT::CoreDD - PRIVATE MLIRQTensorDialect MQT::ProjectWarnings MQT::ProjectOptions) + PRIVATE MQT::ProjectWarnings MQT::ProjectOptions) endif() diff --git a/mlir/lib/Dialect/QTensor/IR/CMakeLists.txt b/mlir/lib/Dialect/QTensor/IR/CMakeLists.txt index 97d7ec241c..3fe82903c0 100644 --- a/mlir/lib/Dialect/QTensor/IR/CMakeLists.txt +++ b/mlir/lib/Dialect/QTensor/IR/CMakeLists.txt @@ -10,6 +10,7 @@ file(GLOB_RECURSE OPERATIONS "${CMAKE_CURRENT_SOURCE_DIR}/Operations/*.cpp") add_mlir_dialect_library( MLIRQTensorDialect + QTensorCanonicalization.cpp QTensorOps.cpp ${OPERATIONS} ADDITIONAL_HEADER_DIRS @@ -18,7 +19,6 @@ add_mlir_dialect_library( MLIRQTensorOpsIncGen LINK_LIBS PRIVATE - MLIRQTensorUtils MLIRIR MLIRDialectUtils MLIRArithDialect diff --git a/mlir/lib/Dialect/QTensor/IR/QTensorCanonicalization.cpp b/mlir/lib/Dialect/QTensor/IR/QTensorCanonicalization.cpp new file mode 100644 index 0000000000..182776456d --- /dev/null +++ b/mlir/lib/Dialect/QTensor/IR/QTensorCanonicalization.cpp @@ -0,0 +1,275 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#include "mlir/Dialect/QCO/IR/QCOOps.h" +#include "mlir/Dialect/QTensor/IR/QTensorDialect.h" +#include "mlir/Dialect/QTensor/IR/QTensorOps.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +using namespace mlir; +using namespace mlir::qtensor; + +namespace { + +struct QTensorAccess { + ExtractOp extract; + InsertOp insert; +}; + +struct BranchQTensorAccesses { + DenseMap accesses; + SmallVector qTensorOperations; +}; + +} // namespace + +/// Analyze a QTensor's complete lifetime in one branch. +/// +/// Supported branches extract distinct constant-index qubits, perform +/// QTensor-independent computation, reinsert one qubit at every extracted +/// index, and yield the resulting QTensor. Dynamic indices, repeated accesses, +/// and partial updates do not match. +static std::optional +analyzeQTensorBranch(Block* block, size_t qTensorArgumentIndex, + size_t qTensorYieldIndex) { + BranchQTensorAccesses result; + Value currentQTensor = block->getArgument(qTensorArgumentIndex); + bool reachedInsertPhase = false; + + while (true) { + assert(currentQTensor.hasOneUse() && "expected linear typing"); + Operation* user = *currentQTensor.getUsers().begin(); + if (user->getBlock() != block) { + return std::nullopt; + } + + if (auto extract = dyn_cast(user)) { + auto index = getConstantIntValue(extract.getIndex()); + if (reachedInsertPhase || !index || + !result.accesses + .try_emplace(*index, QTensorAccess{.extract = extract}) + .second) { + return std::nullopt; + } + result.qTensorOperations.push_back(user); + currentQTensor = extract.getOutTensor(); + continue; + } + + if (auto insert = dyn_cast(user)) { + reachedInsertPhase = true; + auto index = getConstantIntValue(insert.getIndex()); + if (!index) { + return std::nullopt; + } + auto access = result.accesses.find(*index); + if (access == result.accesses.end() || access->second.insert) { + return std::nullopt; + } + access->second.insert = insert; + result.qTensorOperations.push_back(user); + currentQTensor = insert.getResult(); + continue; + } + + auto yield = dyn_cast(user); + if (!yield || user != block->getTerminator() || + qTensorYieldIndex >= yield.getTargets().size() || + yield.getTargets()[qTensorYieldIndex] != currentQTensor || + llvm::any_of(result.accesses, [](const auto& access) { + return !access.second.insert; + })) { + return std::nullopt; + } + return result; + } +} + +/// Move a branch while replacing QTensor accesses with scalar qubits. +static void moveScalarizedQTensorBranch(qco::IfOp oldIf, Block* oldBlock, + Block* newBlock, + size_t qTensorArgumentIndex, + BranchQTensorAccesses& accesses, + ArrayRef indices, + PatternRewriter& rewriter) { + auto oldYield = cast(oldBlock->getTerminator()); + auto scalarArguments = newBlock->getArguments().take_back(indices.size()); + auto carriedArguments = newBlock->getArguments().drop_back(indices.size()); + + SmallVector argumentReplacements; + argumentReplacements.reserve(oldBlock->getNumArguments()); + size_t carriedIndex = 0; + for (size_t oldIndex : llvm::seq(oldBlock->getNumArguments())) { + argumentReplacements.push_back(oldIndex == qTensorArgumentIndex + ? oldIf.getQubits()[qTensorArgumentIndex] + : carriedArguments[carriedIndex++]); + } + assert(carriedIndex == carriedArguments.size()); + rewriter.mergeBlocks(oldBlock, newBlock, argumentReplacements); + + SmallVector scalarYields; + scalarYields.reserve(indices.size()); + for (auto [indexPosition, index] : llvm::enumerate(indices)) { + auto access = accesses.accesses.find(index); + if (access == accesses.accesses.end()) { + scalarYields.push_back(scalarArguments[indexPosition]); + } else { + rewriter.replaceAllUsesWith(access->second.extract.getResult(), + scalarArguments[indexPosition]); + scalarYields.push_back(access->second.insert.getScalar()); + } + } + + auto oldTargets = oldYield.getTargets(); + size_t classicalResultCount = oldIf.getClassicalResults().size(); + SmallVector newYieldValues; + newYieldValues.reserve(oldTargets.size() - 1 + scalarYields.size()); + llvm::append_range(newYieldValues, + oldTargets.take_front(classicalResultCount)); + for (auto [oldIndex, value] : + llvm::enumerate(oldTargets.drop_front(classicalResultCount))) { + if (oldIndex != qTensorArgumentIndex) { + newYieldValues.push_back(value); + } + } + llvm::append_range(newYieldValues, scalarYields); + + rewriter.setInsertionPoint(oldYield); + rewriter.replaceOpWithNewOp(oldYield, newYieldValues); + + for (Operation* operation : llvm::reverse(accesses.qTensorOperations)) { + rewriter.eraseOp(operation); + } +} + +namespace { + +/// Replace constant-index QTensor updates in an if with scalar threading. +/// +/// A QTensor carried through an if hides its qubits from target mapping. This +/// pattern extracts the union of constant indices accessed by either branch, +/// threads those qubits through both branches, and reinserts the results. +/// Untouched elements remain in the QTensor outside the if. +struct ScalarizeQTensorInputs final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(qco::IfOp op, + PatternRewriter& rewriter) const override { + size_t classicalResultCount = op.getClassicalResults().size(); + auto oldQubits = op.getQubits(); + + for (auto [qTensorIndex, qTensor] : llvm::enumerate(oldQubits)) { + auto qTensorType = dyn_cast(qTensor.getType()); + if (!qTensorType || !qTensorType.hasStaticShape()) { + continue; + } + + auto thenAccesses = analyzeQTensorBranch( + op.thenBlock(), qTensorIndex, classicalResultCount + qTensorIndex); + auto elseAccesses = analyzeQTensorBranch( + op.elseBlock(), qTensorIndex, classicalResultCount + qTensorIndex); + if (!thenAccesses || !elseAccesses) { + continue; + } + + SmallVector accessedIndices(thenAccesses->accesses.keys()); + llvm::append_range(accessedIndices, elseAccesses->accesses.keys()); + llvm::sort(accessedIndices); + accessedIndices.erase(llvm::unique(accessedIndices), + accessedIndices.end()); + ArrayRef indices(accessedIndices); + + rewriter.setInsertionPoint(op); + SmallVector indexValues; + SmallVector scalarInputs; + indexValues.reserve(indices.size()); + scalarInputs.reserve(indices.size()); + Value qTensorWithoutScalars = qTensor; + for (int64_t index : indices) { + auto indexValue = + arith::ConstantIndexOp::create(rewriter, op.getLoc(), index); + auto extract = + ExtractOp::create(rewriter, op.getLoc(), qTensorWithoutScalars, + indexValue.getResult()); + indexValues.push_back(indexValue.getResult()); + scalarInputs.push_back(extract.getResult()); + qTensorWithoutScalars = extract.getOutTensor(); + } + + SmallVector newQubits(oldQubits); + newQubits.erase(newQubits.begin() + qTensorIndex); + llvm::append_range(newQubits, scalarInputs); + + auto newIf = qco::IfOp::create( + rewriter, op.getLoc(), op.getClassicalResults().getTypes(), + ValueRange(newQubits).getTypes(), op.getCondition(), newQubits); + newIf->setDiscardableAttrs(op->getDiscardableAttrDictionary()); + + SmallVector locations(newQubits.size(), op.getLoc()); + Block* oldThenBlock = op.thenBlock(); + Block* oldElseBlock = op.elseBlock(); + Block* newThenBlock = + rewriter.createBlock(&newIf.getThenRegion(), {}, + ValueRange(newQubits).getTypes(), locations); + Block* newElseBlock = + rewriter.createBlock(&newIf.getElseRegion(), {}, + ValueRange(newQubits).getTypes(), locations); + moveScalarizedQTensorBranch(op, oldThenBlock, newThenBlock, qTensorIndex, + *thenAccesses, indices, rewriter); + moveScalarizedQTensorBranch(op, oldElseBlock, newElseBlock, qTensorIndex, + *elseAccesses, indices, rewriter); + + rewriter.setInsertionPointAfter(newIf); + Value updatedQTensor = qTensorWithoutScalars; + auto scalarResults = newIf.getLinearResults().take_back(indices.size()); + for (auto [scalar, indexValue] : + llvm::zip_equal(scalarResults, indexValues)) { + updatedQTensor = InsertOp::create(rewriter, op.getLoc(), scalar, + updatedQTensor, indexValue) + .getResult(); + } + + SmallVector replacements( + newIf.getLinearResults().drop_back(indices.size())); + replacements.insert(replacements.begin() + qTensorIndex, updatedQTensor); + replacements.insert(replacements.begin(), + newIf.getClassicalResults().begin(), + newIf.getClassicalResults().end()); + rewriter.replaceOp(op, replacements); + return success(); + } + return failure(); + } +}; +} // namespace + +void QTensorDialect::getCanonicalizationPatterns( + RewritePatternSet& results) const { + results.add(getContext()); +} diff --git a/mlir/lib/Dialect/QTensor/IR/QTensorOps.cpp b/mlir/lib/Dialect/QTensor/IR/QTensorOps.cpp index c26ed12d62..c339198699 100644 --- a/mlir/lib/Dialect/QTensor/IR/QTensorOps.cpp +++ b/mlir/lib/Dialect/QTensor/IR/QTensorOps.cpp @@ -12,6 +12,8 @@ #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" // IWYU pragma: associated +#include + // The following headers are needed for some template instantiations. // IWYU pragma: begin_keep #include diff --git a/mlir/lib/Dialect/QTensor/Utils/CMakeLists.txt b/mlir/lib/Dialect/QTensor/Utils/CMakeLists.txt index 4e3943c1d0..e688756f21 100644 --- a/mlir/lib/Dialect/QTensor/Utils/CMakeLists.txt +++ b/mlir/lib/Dialect/QTensor/Utils/CMakeLists.txt @@ -16,11 +16,12 @@ add_mlir_dialect_library( DEPENDS MLIRQTensorOpsIncGen LINK_LIBS - PUBLIC - MLIRQCODialect PRIVATE + MLIRQTensorDialect MLIRFuncDialect - MLIRSCFDialect) + MLIRSCFDialect + PUBLIC + MLIRQCODialect) mqt_mlir_target_use_project_options(MLIRQTensorUtils) diff --git a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp index c0e9a8a4f8..7cb9c6ebc1 100644 --- a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp +++ b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp @@ -990,343 +990,6 @@ TEST_F(QCOTest, CanonicalizesRedundantClassicalIfResults) { EXPECT_EQ(returnOp.getOperand(2), returnOp.getOperand(1)); } -TEST_F(QCOTest, CanonicalizesConstantIndexQTensorIfToScalarQubits) { - constexpr StringLiteral mlirCode = R"mlir( - module { - func.func @main(%condition: i1) -> i1 { - %c0 = arith.constant 0 : index - %c1 = arith.constant 1 : index - %c2 = arith.constant 2 : index - %c3 = arith.constant 3 : index - %tensor0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - %flag, %tensor1 = qco.if %condition - args(%arg0 = %tensor0) -> (i1, tensor<3x!qco.qubit>) { - %tensor2, %q0 = qtensor.extract %arg0[%c0] - : tensor<3x!qco.qubit> - %tensor3, %q1 = qtensor.extract %tensor2[%c1] - : tensor<3x!qco.qubit> - %q2, %q3 = qco.swap %q0, %q1 - : !qco.qubit, !qco.qubit -> !qco.qubit, !qco.qubit - %tensor4 = qtensor.insert %q3 into %tensor3[%c1] - : tensor<3x!qco.qubit> - %tensor5 = qtensor.insert %q2 into %tensor4[%c0] - : tensor<3x!qco.qubit> - %true = arith.constant true - qco.yield %true, %tensor5 : i1, tensor<3x!qco.qubit> - } else args(%arg0 = %tensor0) { - %tensor2, %q0 = qtensor.extract %arg0[%c2] - : tensor<3x!qco.qubit> - %q1 = qco.z %q0 : !qco.qubit -> !qco.qubit - %tensor3 = qtensor.insert %q1 into %tensor2[%c2] - : tensor<3x!qco.qubit> - %false = arith.constant false - qco.yield %false, %tensor3 : i1, tensor<3x!qco.qubit> - } {test.marker = "preserved"} - qtensor.dealloc %tensor1 : tensor<3x!qco.qubit> - return %flag : i1 - } - } - )mlir"; - - auto moduleOp = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - IfOp ifOp; - moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); - ASSERT_TRUE(ifOp); - ASSERT_EQ(ifOp.getClassicalResults().size(), 1); - ASSERT_EQ(ifOp.getQubits().size(), 3); - ASSERT_EQ(ifOp.getLinearResults().size(), 3); - EXPECT_EQ( - cast(ifOp->getDiscardableAttr("test.marker")).getValue(), - "preserved"); - EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { - return isa(value.getType()); - })); - EXPECT_TRUE(llvm::all_of(ifOp.getLinearResults(), [](Value value) { - return isa(value.getType()); - })); - - size_t nestedExtracts = 0; - size_t nestedInserts = 0; - size_t swaps = 0; - size_t zs = 0; - ifOp->walk([&](Operation* operation) { - nestedExtracts += isa(operation); - nestedInserts += isa(operation); - swaps += isa(operation); - zs += isa(operation); - }); - EXPECT_EQ(nestedExtracts, 0); - EXPECT_EQ(nestedInserts, 0); - EXPECT_EQ(swaps, 1); - EXPECT_EQ(zs, 1); - - size_t extracts = 0; - size_t inserts = 0; - moduleOp->walk([&](qtensor::ExtractOp) { ++extracts; }); - moduleOp->walk([&](qtensor::InsertOp) { ++inserts; }); - EXPECT_EQ(extracts, 3); - EXPECT_EQ(inserts, 3); -} - -TEST_F(QCOTest, ScalarizesOnlyAccessedQTensorElements) { - constexpr StringLiteral mlirCode = R"mlir( - module { - func.func @main(%condition: i1) { - %c1 = arith.constant 1 : index - %c3 = arith.constant 3 : index - %tensor0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - %tensor1 = qco.if %condition - args(%arg0 = %tensor0) -> (tensor<3x!qco.qubit>) { - %tensor2, %q0 = qtensor.extract %arg0[%c1] - : tensor<3x!qco.qubit> - %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit - %tensor3 = qtensor.insert %q1 into %tensor2[%c1] - : tensor<3x!qco.qubit> - qco.yield %tensor3 : tensor<3x!qco.qubit> - } else args(%arg0 = %tensor0) { - qco.yield %arg0 : tensor<3x!qco.qubit> - } - qtensor.dealloc %tensor1 : tensor<3x!qco.qubit> - return - } - } - )mlir"; - - auto moduleOp = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - IfOp ifOp; - moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); - ASSERT_TRUE(ifOp); - ASSERT_EQ(ifOp.getQubits().size(), 1); - ASSERT_EQ(ifOp.getLinearResults().size(), 1); - EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { - return isa(value.getType()); - })); - - auto thenValues = ifOp.thenYield().getTargets(); - auto elseValues = ifOp.elseYield().getTargets(); - ASSERT_EQ(thenValues.size(), 1); - ASSERT_EQ(elseValues.size(), 1); - EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); - EXPECT_EQ(elseValues[0], ifOp.elseBlock()->getArgument(0)); - - size_t extracts = 0; - size_t inserts = 0; - moduleOp->walk([&](qtensor::ExtractOp) { ++extracts; }); - moduleOp->walk([&](qtensor::InsertOp) { ++inserts; }); - EXPECT_EQ(extracts, 1); - EXPECT_EQ(inserts, 1); -} - -TEST_F(QCOTest, ForwardsUnaccessedQTensorAroundIf) { - constexpr StringLiteral mlirCode = R"mlir( - module { - func.func @main(%condition: i1) { - %c2 = arith.constant 2 : index - %tensor0 = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %q0 = qco.alloc : !qco.qubit - %tensor1, %q1 = qco.if %condition - args(%tensor = %tensor0, %q = %q0) - -> (tensor<2x!qco.qubit>, !qco.qubit) { - %q2 = qco.h %q : !qco.qubit -> !qco.qubit - qco.yield %tensor, %q2 : tensor<2x!qco.qubit>, !qco.qubit - } else args(%tensor = %tensor0, %q = %q0) { - qco.yield %tensor, %q : tensor<2x!qco.qubit>, !qco.qubit - } - qtensor.dealloc %tensor1 : tensor<2x!qco.qubit> - qco.sink %q1 : !qco.qubit - return - } - } - )mlir"; - - auto moduleOp = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - IfOp ifOp; - moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); - ASSERT_TRUE(ifOp); - ASSERT_EQ(ifOp.getQubits().size(), 1); - ASSERT_EQ(ifOp.getLinearResults().size(), 1); - EXPECT_TRUE(isa(ifOp.getQubits()[0].getType())); - EXPECT_TRUE(isa(ifOp.getLinearResults()[0].getType())); - - auto thenValues = ifOp.thenYield().getTargets(); - auto elseValues = ifOp.elseYield().getTargets(); - ASSERT_EQ(thenValues.size(), 1); - ASSERT_EQ(elseValues.size(), 1); - EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); - for (auto [index, value] : llvm::enumerate(elseValues)) { - EXPECT_EQ(value, ifOp.elseBlock()->getArgument(index)); - } -} - -TEST_F(QCOTest, PreservesInterleavedResultOrderWhenScalarizingQTensors) { - constexpr StringLiteral mlirCode = R"mlir( - module { - func.func @main(%condition: i1) { - %c0 = arith.constant 0 : index - %c1 = arith.constant 1 : index - %tensorA0 = qtensor.alloc(%c1) : tensor<1x!qco.qubit> - %middle0 = qco.alloc : !qco.qubit - %tensorB0 = qtensor.alloc(%c1) : tensor<1x!qco.qubit> - %tensorA1, %middle1, %tensorB1 = - qco.if %condition - args(%tensorA = %tensorA0, %middle = %middle0, - %tensorB = %tensorB0) - -> (tensor<1x!qco.qubit>, !qco.qubit, - tensor<1x!qco.qubit>) { - %tensorA2, %tensorAQubit = qtensor.extract %tensorA[%c0] - : tensor<1x!qco.qubit> - %tensorAQubitOut = qco.x %tensorAQubit - : !qco.qubit -> !qco.qubit - %tensorA3 = qtensor.insert %tensorAQubitOut into %tensorA2[%c0] - : tensor<1x!qco.qubit> - %middleOut = qco.y %middle : !qco.qubit -> !qco.qubit - %tensorB2, %tensorBQubit = qtensor.extract %tensorB[%c0] - : tensor<1x!qco.qubit> - %tensorBQubitOut = qco.z %tensorBQubit - : !qco.qubit -> !qco.qubit - %tensorB3 = qtensor.insert %tensorBQubitOut into %tensorB2[%c0] - : tensor<1x!qco.qubit> - qco.yield %tensorA3, %middleOut, %tensorB3 - : tensor<1x!qco.qubit>, !qco.qubit, tensor<1x!qco.qubit> - } else args(%tensorA = %tensorA0, %middle = %middle0, - %tensorB = %tensorB0) { - qco.yield %tensorA, %middle, %tensorB - : tensor<1x!qco.qubit>, !qco.qubit, tensor<1x!qco.qubit> - } - %middle2 = qco.t %middle1 : !qco.qubit -> !qco.qubit - qtensor.dealloc %tensorA1 : tensor<1x!qco.qubit> - qco.sink %middle2 : !qco.qubit - qtensor.dealloc %tensorB1 : tensor<1x!qco.qubit> - return - } - } - )mlir"; - - auto moduleOp = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - IfOp ifOp; - TOp postMiddle; - SmallVector insertedScalars; - moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); - moduleOp->walk([&](TOp candidate) { postMiddle = candidate; }); - moduleOp->walk([&](qtensor::InsertOp insert) { - insertedScalars.push_back(insert.getScalar()); - }); - - ASSERT_TRUE(ifOp); - ASSERT_EQ(ifOp.getQubits().size(), 3); - ASSERT_EQ(ifOp.getLinearResults().size(), 3); - EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { - return isa(value.getType()); - })); - EXPECT_TRUE(llvm::all_of(ifOp.getLinearResults(), [](Value value) { - return isa(value.getType()); - })); - - ASSERT_TRUE(postMiddle); - EXPECT_EQ(cast(postMiddle.getOperation()) - .getInputQubits() - .front(), - ifOp.getLinearResults()[0]); - - ASSERT_EQ(insertedScalars.size(), 2); - EXPECT_TRUE(llvm::is_contained(insertedScalars, ifOp.getLinearResults()[1])); - EXPECT_TRUE(llvm::is_contained(insertedScalars, ifOp.getLinearResults()[2])); - - auto thenValues = ifOp.thenYield().getTargets(); - ASSERT_EQ(thenValues.size(), 3); - EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); - EXPECT_TRUE(isa(thenValues[1].getDefiningOp())); - EXPECT_TRUE(isa(thenValues[2].getDefiningOp())); -} - -TEST_F(QCOTest, LeavesUnsupportedQTensorIfUnchanged) { - constexpr std::array mlirCodes = { - R"mlir( - module { - func.func @main(%condition: i1, %index: index) { - %c2 = arith.constant 2 : index - %tensor0 = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %tensor1 = qco.if %condition - args(%arg0 = %tensor0) -> (tensor<2x!qco.qubit>) { - %tensor2, %q0 = qtensor.extract %arg0[%index] - : tensor<2x!qco.qubit> - %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit - %tensor3 = qtensor.insert %q1 into %tensor2[%index] - : tensor<2x!qco.qubit> - qco.yield %tensor3 : tensor<2x!qco.qubit> - } else args(%arg0 = %tensor0) { - qco.yield %arg0 : tensor<2x!qco.qubit> - } - qtensor.dealloc %tensor1 : tensor<2x!qco.qubit> - return - } - } - )mlir", - R"mlir( - module { - func.func @main(%condition: i1, %size: index) { - %c0 = arith.constant 0 : index - %tensor0 = qtensor.alloc(%size) : tensor - %tensor1 = qco.if %condition - args(%arg0 = %tensor0) -> (tensor) { - %tensor2, %q0 = qtensor.extract %arg0[%c0] - : tensor - %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit - %tensor3 = qtensor.insert %q1 into %tensor2[%c0] - : tensor - qco.yield %tensor3 : tensor - } else args(%arg0 = %tensor0) { - qco.yield %arg0 : tensor - } - qtensor.dealloc %tensor1 : tensor - return - } - } - )mlir"}; - - for (StringRef mlirCode : mlirCodes) { - auto moduleOp = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(moduleOp); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); - ASSERT_TRUE(succeeded(verify(*moduleOp))); - - IfOp ifOp; - moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); - ASSERT_TRUE(ifOp); - ASSERT_EQ(ifOp.getQubits().size(), 1); - EXPECT_TRUE(isa(ifOp.getQubits().front().getType())); - size_t nestedExtracts = 0; - size_t nestedInserts = 0; - ifOp->walk([&](Operation* operation) { - nestedExtracts += isa(operation); - nestedInserts += isa(operation); - }); - EXPECT_EQ(nestedExtracts, 1); - EXPECT_EQ(nestedInserts, 1); - } -} - TEST_F(QCOTest, IndexSwitchParser) { // Test IndexSwitch parser const char* mlirCode = R"( diff --git a/mlir/unittests/Dialect/QIR/Execution/Runtime/CMakeLists.txt b/mlir/unittests/Dialect/QIR/Execution/Runtime/CMakeLists.txt index 2835351439..a0c5c94a11 100644 --- a/mlir/unittests/Dialect/QIR/Execution/Runtime/CMakeLists.txt +++ b/mlir/unittests/Dialect/QIR/Execution/Runtime/CMakeLists.txt @@ -30,6 +30,7 @@ macro(ADD_QIR_CIRCUIT target_name circuit_path) DEPENDS ${circuit_path} COMMENT "Compiling ${circuit_path} to ${circuit_name}.o with llc") add_executable(${target_name} ${circuit_name}.o) + set_property(TARGET ${target_name} PROPERTY BUILD_WITH_INSTALL_RPATH FALSE) target_link_libraries(${target_name} PRIVATE MQT::CoreQIRRuntime) set_target_properties(${target_name} PROPERTIES LINKER_LANGUAGE CXX) endif() diff --git a/mlir/unittests/Dialect/QTensor/IR/CMakeLists.txt b/mlir/unittests/Dialect/QTensor/IR/CMakeLists.txt index b7e2227682..f1ae02b832 100644 --- a/mlir/unittests/Dialect/QTensor/IR/CMakeLists.txt +++ b/mlir/unittests/Dialect/QTensor/IR/CMakeLists.txt @@ -7,7 +7,7 @@ # Licensed under the MIT License set(qtensor_ir_target mqt-core-mlir-unittest-qtensor-ir) -add_executable(${qtensor_ir_target} test_qtensor_ir.cpp) +add_executable(${qtensor_ir_target} test_qtensor_canonicalization.cpp test_qtensor_ir.cpp) target_link_libraries(${qtensor_ir_target} PRIVATE GTest::gtest_main MLIRParser MLIRSupportMQT MLIRQCOProgramBuilder MLIRQCOPrograms) mqt_mlir_configure_unittest_target(${qtensor_ir_target}) diff --git a/mlir/unittests/Dialect/QTensor/IR/test_qtensor_canonicalization.cpp b/mlir/unittests/Dialect/QTensor/IR/test_qtensor_canonicalization.cpp new file mode 100644 index 0000000000..0468f15f26 --- /dev/null +++ b/mlir/unittests/Dialect/QTensor/IR/test_qtensor_canonicalization.cpp @@ -0,0 +1,392 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#include "mlir/Dialect/QCO/IR/QCOInterfaces.h" +#include "mlir/Dialect/QCO/IR/QCOOps.h" +#include "mlir/Dialect/QTensor/IR/QTensorDialect.h" +#include "mlir/Dialect/QTensor/IR/QTensorOps.h" +#include "mlir/Support/Passes.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +using namespace mlir; +using namespace mlir::qco; + +namespace { + +class QTensorCanonicalizationTest : public testing::Test { +protected: + MLIRContext context_; + + void SetUp() override { + context_.loadDialect(); + } +}; + +TEST_F(QTensorCanonicalizationTest, + CanonicalizesConstantIndexQTensorIfToScalarQubits) { + constexpr StringLiteral mlirCode = R"mlir( + module { + func.func @main(%condition: i1) -> i1 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c2 = arith.constant 2 : index + %c3 = arith.constant 3 : index + %tensor0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> + %flag, %tensor1 = qco.if %condition + args(%arg0 = %tensor0) -> (i1, tensor<3x!qco.qubit>) { + %tensor2, %q0 = qtensor.extract %arg0[%c0] + : tensor<3x!qco.qubit> + %tensor3, %q1 = qtensor.extract %tensor2[%c1] + : tensor<3x!qco.qubit> + %q2, %q3 = qco.swap %q0, %q1 + : !qco.qubit, !qco.qubit -> !qco.qubit, !qco.qubit + %tensor4 = qtensor.insert %q3 into %tensor3[%c1] + : tensor<3x!qco.qubit> + %tensor5 = qtensor.insert %q2 into %tensor4[%c0] + : tensor<3x!qco.qubit> + %true = arith.constant true + qco.yield %true, %tensor5 : i1, tensor<3x!qco.qubit> + } else args(%arg0 = %tensor0) { + %tensor2, %q0 = qtensor.extract %arg0[%c2] + : tensor<3x!qco.qubit> + %q1 = qco.z %q0 : !qco.qubit -> !qco.qubit + %tensor3 = qtensor.insert %q1 into %tensor2[%c2] + : tensor<3x!qco.qubit> + %false = arith.constant false + qco.yield %false, %tensor3 : i1, tensor<3x!qco.qubit> + } {test.marker = "preserved"} + qtensor.dealloc %tensor1 : tensor<3x!qco.qubit> + return %flag : i1 + } + } + )mlir"; + + auto moduleOp = parseSourceString(mlirCode, &context_); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + PassManager pm(&context_); + pm.addPass(createCanonicalizerPass()); + ASSERT_TRUE(succeeded(pm.run(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + IfOp ifOp; + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); + ASSERT_TRUE(ifOp); + ASSERT_EQ(ifOp.getClassicalResults().size(), 1); + ASSERT_EQ(ifOp.getQubits().size(), 3); + ASSERT_EQ(ifOp.getLinearResults().size(), 3); + EXPECT_EQ( + cast(ifOp->getDiscardableAttr("test.marker")).getValue(), + "preserved"); + EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { + return isa(value.getType()); + })); + EXPECT_TRUE(llvm::all_of(ifOp.getLinearResults(), [](Value value) { + return isa(value.getType()); + })); + + size_t nestedExtracts = 0; + size_t nestedInserts = 0; + size_t swaps = 0; + size_t zs = 0; + ifOp->walk([&](Operation* operation) { + nestedExtracts += isa(operation); + nestedInserts += isa(operation); + swaps += isa(operation); + zs += isa(operation); + }); + EXPECT_EQ(nestedExtracts, 0); + EXPECT_EQ(nestedInserts, 0); + EXPECT_EQ(swaps, 1); + EXPECT_EQ(zs, 1); + + size_t extracts = 0; + size_t inserts = 0; + moduleOp->walk([&](qtensor::ExtractOp) { ++extracts; }); + moduleOp->walk([&](qtensor::InsertOp) { ++inserts; }); + EXPECT_EQ(extracts, 3); + EXPECT_EQ(inserts, 3); +} + +TEST_F(QTensorCanonicalizationTest, ScalarizesOnlyAccessedQTensorElements) { + constexpr StringLiteral mlirCode = R"mlir( + module { + func.func @main(%condition: i1) { + %c1 = arith.constant 1 : index + %c3 = arith.constant 3 : index + %tensor0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> + %tensor1 = qco.if %condition + args(%arg0 = %tensor0) -> (tensor<3x!qco.qubit>) { + %tensor2, %q0 = qtensor.extract %arg0[%c1] + : tensor<3x!qco.qubit> + %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit + %tensor3 = qtensor.insert %q1 into %tensor2[%c1] + : tensor<3x!qco.qubit> + qco.yield %tensor3 : tensor<3x!qco.qubit> + } else args(%arg0 = %tensor0) { + qco.yield %arg0 : tensor<3x!qco.qubit> + } + qtensor.dealloc %tensor1 : tensor<3x!qco.qubit> + return + } + } + )mlir"; + + auto moduleOp = parseSourceString(mlirCode, &context_); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + IfOp ifOp; + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); + ASSERT_TRUE(ifOp); + ASSERT_EQ(ifOp.getQubits().size(), 1); + ASSERT_EQ(ifOp.getLinearResults().size(), 1); + EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { + return isa(value.getType()); + })); + + auto thenValues = ifOp.thenYield().getTargets(); + auto elseValues = ifOp.elseYield().getTargets(); + ASSERT_EQ(thenValues.size(), 1); + ASSERT_EQ(elseValues.size(), 1); + EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); + EXPECT_EQ(elseValues[0], ifOp.elseBlock()->getArgument(0)); + + size_t extracts = 0; + size_t inserts = 0; + moduleOp->walk([&](qtensor::ExtractOp) { ++extracts; }); + moduleOp->walk([&](qtensor::InsertOp) { ++inserts; }); + EXPECT_EQ(extracts, 1); + EXPECT_EQ(inserts, 1); +} + +TEST_F(QTensorCanonicalizationTest, ForwardsUnaccessedQTensorAroundIf) { + constexpr StringLiteral mlirCode = R"mlir( + module { + func.func @main(%condition: i1) { + %c2 = arith.constant 2 : index + %tensor0 = qtensor.alloc(%c2) : tensor<2x!qco.qubit> + %q0 = qco.alloc : !qco.qubit + %tensor1, %q1 = qco.if %condition + args(%tensor = %tensor0, %q = %q0) + -> (tensor<2x!qco.qubit>, !qco.qubit) { + %q2 = qco.h %q : !qco.qubit -> !qco.qubit + qco.yield %tensor, %q2 : tensor<2x!qco.qubit>, !qco.qubit + } else args(%tensor = %tensor0, %q = %q0) { + qco.yield %tensor, %q : tensor<2x!qco.qubit>, !qco.qubit + } + qtensor.dealloc %tensor1 : tensor<2x!qco.qubit> + qco.sink %q1 : !qco.qubit + return + } + } + )mlir"; + + auto moduleOp = parseSourceString(mlirCode, &context_); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + IfOp ifOp; + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); + ASSERT_TRUE(ifOp); + ASSERT_EQ(ifOp.getQubits().size(), 1); + ASSERT_EQ(ifOp.getLinearResults().size(), 1); + EXPECT_TRUE(isa(ifOp.getQubits()[0].getType())); + EXPECT_TRUE(isa(ifOp.getLinearResults()[0].getType())); + + auto thenValues = ifOp.thenYield().getTargets(); + auto elseValues = ifOp.elseYield().getTargets(); + ASSERT_EQ(thenValues.size(), 1); + ASSERT_EQ(elseValues.size(), 1); + EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); + for (auto [index, value] : llvm::enumerate(elseValues)) { + EXPECT_EQ(value, ifOp.elseBlock()->getArgument(index)); + } +} + +TEST_F(QTensorCanonicalizationTest, + PreservesInterleavedResultOrderWhenScalarizingQTensors) { + constexpr StringLiteral mlirCode = R"mlir( + module { + func.func @main(%condition: i1) { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %tensorA0 = qtensor.alloc(%c1) : tensor<1x!qco.qubit> + %middle0 = qco.alloc : !qco.qubit + %tensorB0 = qtensor.alloc(%c1) : tensor<1x!qco.qubit> + %tensorA1, %middle1, %tensorB1 = + qco.if %condition + args(%tensorA = %tensorA0, %middle = %middle0, + %tensorB = %tensorB0) + -> (tensor<1x!qco.qubit>, !qco.qubit, + tensor<1x!qco.qubit>) { + %tensorA2, %tensorAQubit = qtensor.extract %tensorA[%c0] + : tensor<1x!qco.qubit> + %tensorAQubitOut = qco.x %tensorAQubit + : !qco.qubit -> !qco.qubit + %tensorA3 = qtensor.insert %tensorAQubitOut into %tensorA2[%c0] + : tensor<1x!qco.qubit> + %middleOut = qco.y %middle : !qco.qubit -> !qco.qubit + %tensorB2, %tensorBQubit = qtensor.extract %tensorB[%c0] + : tensor<1x!qco.qubit> + %tensorBQubitOut = qco.z %tensorBQubit + : !qco.qubit -> !qco.qubit + %tensorB3 = qtensor.insert %tensorBQubitOut into %tensorB2[%c0] + : tensor<1x!qco.qubit> + qco.yield %tensorA3, %middleOut, %tensorB3 + : tensor<1x!qco.qubit>, !qco.qubit, tensor<1x!qco.qubit> + } else args(%tensorA = %tensorA0, %middle = %middle0, + %tensorB = %tensorB0) { + qco.yield %tensorA, %middle, %tensorB + : tensor<1x!qco.qubit>, !qco.qubit, tensor<1x!qco.qubit> + } + %middle2 = qco.t %middle1 : !qco.qubit -> !qco.qubit + qtensor.dealloc %tensorA1 : tensor<1x!qco.qubit> + qco.sink %middle2 : !qco.qubit + qtensor.dealloc %tensorB1 : tensor<1x!qco.qubit> + return + } + } + )mlir"; + + auto moduleOp = parseSourceString(mlirCode, &context_); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + IfOp ifOp; + TOp postMiddle; + SmallVector insertedScalars; + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](TOp candidate) { postMiddle = candidate; }); + moduleOp->walk([&](qtensor::InsertOp insert) { + insertedScalars.push_back(insert.getScalar()); + }); + + ASSERT_TRUE(ifOp); + ASSERT_EQ(ifOp.getQubits().size(), 3); + ASSERT_EQ(ifOp.getLinearResults().size(), 3); + EXPECT_TRUE(llvm::all_of(ifOp.getQubits(), [](Value value) { + return isa(value.getType()); + })); + EXPECT_TRUE(llvm::all_of(ifOp.getLinearResults(), [](Value value) { + return isa(value.getType()); + })); + + ASSERT_TRUE(postMiddle); + EXPECT_EQ(cast(postMiddle.getOperation()) + .getInputQubits() + .front(), + ifOp.getLinearResults()[0]); + + ASSERT_EQ(insertedScalars.size(), 2); + EXPECT_TRUE(llvm::is_contained(insertedScalars, ifOp.getLinearResults()[1])); + EXPECT_TRUE(llvm::is_contained(insertedScalars, ifOp.getLinearResults()[2])); + + auto thenValues = ifOp.thenYield().getTargets(); + ASSERT_EQ(thenValues.size(), 3); + EXPECT_TRUE(isa(thenValues[0].getDefiningOp())); + EXPECT_TRUE(isa(thenValues[1].getDefiningOp())); + EXPECT_TRUE(isa(thenValues[2].getDefiningOp())); +} + +TEST_F(QTensorCanonicalizationTest, LeavesUnsupportedQTensorIfUnchanged) { + constexpr std::array mlirCodes = { + R"mlir( + module { + func.func @main(%condition: i1, %index: index) { + %c2 = arith.constant 2 : index + %tensor0 = qtensor.alloc(%c2) : tensor<2x!qco.qubit> + %tensor1 = qco.if %condition + args(%arg0 = %tensor0) -> (tensor<2x!qco.qubit>) { + %tensor2, %q0 = qtensor.extract %arg0[%index] + : tensor<2x!qco.qubit> + %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit + %tensor3 = qtensor.insert %q1 into %tensor2[%index] + : tensor<2x!qco.qubit> + qco.yield %tensor3 : tensor<2x!qco.qubit> + } else args(%arg0 = %tensor0) { + qco.yield %arg0 : tensor<2x!qco.qubit> + } + qtensor.dealloc %tensor1 : tensor<2x!qco.qubit> + return + } + } + )mlir", + R"mlir( + module { + func.func @main(%condition: i1, %size: index) { + %c0 = arith.constant 0 : index + %tensor0 = qtensor.alloc(%size) : tensor + %tensor1 = qco.if %condition + args(%arg0 = %tensor0) -> (tensor) { + %tensor2, %q0 = qtensor.extract %arg0[%c0] + : tensor + %q1 = qco.x %q0 : !qco.qubit -> !qco.qubit + %tensor3 = qtensor.insert %q1 into %tensor2[%c0] + : tensor + qco.yield %tensor3 : tensor + } else args(%arg0 = %tensor0) { + qco.yield %arg0 : tensor + } + qtensor.dealloc %tensor1 : tensor + return + } + } + )mlir"}; + + for (StringRef mlirCode : mlirCodes) { + auto moduleOp = parseSourceString(mlirCode, &context_); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + IfOp ifOp; + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); + ASSERT_TRUE(ifOp); + ASSERT_EQ(ifOp.getQubits().size(), 1); + EXPECT_TRUE(isa(ifOp.getQubits().front().getType())); + size_t nestedExtracts = 0; + size_t nestedInserts = 0; + ifOp->walk([&](Operation* operation) { + nestedExtracts += isa(operation); + nestedInserts += isa(operation); + }); + EXPECT_EQ(nestedExtracts, 1); + EXPECT_EQ(nestedInserts, 1); + } +} + +} // namespace