diff --git a/mlir/lib/Dialect/QCO/Transforms/Optimizations/HadamardLifting.cpp b/mlir/lib/Dialect/QCO/Transforms/Optimizations/HadamardLifting.cpp index 49ac6f7a9b..0a24b7c621 100644 --- a/mlir/lib/Dialect/QCO/Transforms/Optimizations/HadamardLifting.cpp +++ b/mlir/lib/Dialect/QCO/Transforms/Optimizations/HadamardLifting.cpp @@ -12,6 +12,7 @@ #include "mlir/Dialect/MQT/Utils/Modifiers.h" #include "mlir/Dialect/QCO/IR/QCOInterfaces.h" #include "mlir/Dialect/QCO/IR/QCOOps.h" +#include "mlir/Dialect/QCO/QCOUtils.h" #include "mlir/Dialect/QCO/Transforms/Passes.h" #include @@ -228,6 +229,10 @@ struct HadamardLifting final : impl::HadamardLiftingBase { void runOnOperation() override { auto op = getOperation(); auto* ctx = &getContext(); + if (failed(qco::verifyLinearity(op))) { + signalPassFailure(); + return; + } // Define the set of patterns to use. RewritePatternSet patterns(ctx); diff --git a/mlir/unittests/Dialect/QCO/Transforms/Optimizations/test_qco_hadamard_lifting.cpp b/mlir/unittests/Dialect/QCO/Transforms/Optimizations/test_qco_hadamard_lifting.cpp index 56e38eff1c..214cf35ce8 100644 --- a/mlir/unittests/Dialect/QCO/Transforms/Optimizations/test_qco_hadamard_lifting.cpp +++ b/mlir/unittests/Dialect/QCO/Transforms/Optimizations/test_qco_hadamard_lifting.cpp @@ -20,8 +20,11 @@ #include #include #include +#include #include #include +#include +#include #include #include #include @@ -81,6 +84,56 @@ class QCOHadamardLiftingTest : public testing::Test { } // namespace +TEST_F(QCOHadamardLiftingTest, HandlesUnusedPauliOutput) { + auto input = parseSourceString(R"mlir( + module { + func.func @main() { + %q = qco.static 0 : !qco.qubit + %unused = qco.x %q : !qco.qubit -> !qco.qubit + return + } + } + )mlir", + &context); + ASSERT_TRUE(input); + ASSERT_TRUE(succeeded(verify(*input))); + OwningOpRef original(input->clone()); + + EXPECT_TRUE(failed(runHadamardLiftingPass(*input))); + EXPECT_TRUE(OperationEquivalence::isEquivalentTo( + input->getOperation(), original->getOperation(), + OperationEquivalence::Flags::None)); +} + +TEST_F(QCOHadamardLiftingTest, HandlesUnusedCnotControlOutput) { + auto input = parseSourceString(R"mlir( + module { + func.func @main() { + %control = qco.static 0 : !qco.qubit + %target = qco.static 1 : !qco.qubit + %unused, %target_out = qco.ctrl(%control) targets(%arg = %target) { + %body = qco.x %arg : !qco.qubit -> !qco.qubit + qco.yield %body : !qco.qubit + } : ({!qco.qubit}, {!qco.qubit}) + -> ({!qco.qubit}, {!qco.qubit}) + %hadamard = qco.h %target_out : !qco.qubit -> !qco.qubit + %measured, %result = qco.measure %hadamard : !qco.qubit + qco.sink %measured : !qco.qubit + return + } + } + )mlir", + &context); + ASSERT_TRUE(input); + ASSERT_TRUE(succeeded(verify(*input))); + OwningOpRef original(input->clone()); + + EXPECT_TRUE(failed(runHadamardLiftingPass(*input))); + EXPECT_TRUE(OperationEquivalence::isEquivalentTo( + input->getOperation(), original->getOperation(), + OperationEquivalence::Flags::None)); +} + // ################################################## // # Raise Hadamard over uncontrolled Pauli gate Tests // ##################################################