| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| #include "QuantumDialect.h" |
| #include "QuantumOps.h" |
| #include "mlir/Dialect/Arith/IR/Arith.h" |
| #include "mlir/IR/PatternMatch.h" |
| #include "mlir/Transforms/GreedyPatternRewriteDriver.h" |
|
|
| using namespace mlir; |
| using namespace mlir::quantum; |
|
|
| |
| |
| |
| static bool isHadamard(UnitaryOp op) { |
| if (op.getQubits().size() != 1) |
| return false; |
| if (op.getAxis() && *op.getAxis() != "Y") |
| return false; |
|
|
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return false; |
|
|
| |
| auto angle = angles[0].dyn_cast<FloatAttr>(); |
| if (!angle) |
| return false; |
|
|
| return std::abs(angle.getValueAsDouble() - 0.5) < 1e-10; |
| } |
|
|
| |
| |
| |
| static bool isTGate(UnitaryOp op) { |
| if (op.getQubits().size() != 1) |
| return false; |
|
|
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return false; |
|
|
| auto angle = angles[0].dyn_cast<FloatAttr>(); |
| if (!angle) |
| return false; |
|
|
| |
| return std::abs(angle.getValueAsDouble() - 0.25) < 1e-10; |
| } |
|
|
| |
| |
| |
| static bool isSGate(UnitaryOp op) { |
| if (op.getQubits().size() != 1) |
| return false; |
|
|
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return false; |
|
|
| auto angle = angles[0].dyn_cast<FloatAttr>(); |
| if (!angle) |
| return false; |
|
|
| |
| return std::abs(angle.getValueAsDouble() - 0.5) < 1e-10; |
| } |
|
|
| |
| |
| |
| struct HHCancellation : public OpRewritePattern<UnitaryOp> { |
| using OpRewritePattern::OpRewritePattern; |
|
|
| LogicalResult matchAndRewrite(UnitaryOp op, |
| PatternRewriter &rewriter) const override { |
| if (!isHadamard(op)) |
| return failure(); |
|
|
| |
| Value qubit = op.getQubits()[0]; |
| auto prevOp = qubit.getDefiningOp<UnitaryOp>(); |
| if (!prevOp || !isHadamard(prevOp)) |
| return failure(); |
|
|
| |
| if (prevOp.getQubits()[0] != qubit) |
| return failure(); |
|
|
| |
| rewriter.replaceOp(op, prevOp.getQubits()); |
| return success(); |
| } |
| }; |
|
|
| |
| |
| |
| struct TripleTCancellation : public OpRewritePattern<UnitaryOp> { |
| using OpRewritePattern::OpRewritePattern; |
|
|
| LogicalResult matchAndRewrite(UnitaryOp op, |
| PatternRewriter &rewriter) const override { |
| if (!isTGate(op)) |
| return failure(); |
|
|
| |
| Value qubit = op.getQubits()[0]; |
| auto prev1 = qubit.getDefiningOp<UnitaryOp>(); |
| if (!prev1 || !isTGate(prev1)) |
| return failure(); |
| if (prev1.getQubits()[0] != qubit) |
| return failure(); |
|
|
| Value qubit1 = prev1.getQubits()[0]; |
| auto prev2 = qubit1.getDefiningOp<UnitaryOp>(); |
| if (!prev2 || !isTGate(prev2)) |
| return failure(); |
| if (prev2.getQubits()[0] != qubit1) |
| return failure(); |
|
|
| |
| |
| auto loc = op.getLoc(); |
| auto sAngle = rewriter.getFloatAttr(rewriter.getF64Type(), 0.5); |
| auto sAngles = rewriter.getArrayAttr({sAngle}); |
|
|
| |
| Value q0 = prev2.getQubits()[0]; |
| auto s1 = rewriter.create<UnitaryOp>( |
| loc, TypeRange{q0.getType()}, sAngles, StringAttr{}, |
| ValueRange{q0}); |
|
|
| |
| auto s2 = rewriter.create<UnitaryOp>( |
| loc, TypeRange{q0.getType()}, sAngles, StringAttr{}, |
| s1.getResults()); |
|
|
| rewriter.replaceOp(op, s2.getResults()); |
| return success(); |
| } |
| }; |
|
|
| |
| |
| |
| struct IdentityElimination : public OpRewritePattern<UnitaryOp> { |
| using OpRewritePattern::OpRewritePattern; |
|
|
| LogicalResult matchAndRewrite(UnitaryOp op, |
| PatternRewriter &rewriter) const override { |
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return failure(); |
|
|
| auto angle = angles[0].dyn_cast<FloatAttr>(); |
| if (!angle) |
| return failure(); |
|
|
| |
| if (std::abs(angle.getValueAsDouble()) > 1e-10) |
| return failure(); |
|
|
| |
| rewriter.replaceOp(op, op.getQubits()); |
| return success(); |
| } |
| }; |
|
|
| |
| |
| |
| struct RzCancellation : public OpRewritePattern<UnitaryOp> { |
| using OpRewritePattern::OpRewritePattern; |
|
|
| LogicalResult matchAndRewrite(UnitaryOp op, |
| PatternRewriter &rewriter) const override { |
| |
| if (op.getQubits().size() != 1) |
| return failure(); |
| if (op.getAxis() && *op.getAxis() != "Z") |
| return failure(); |
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return failure(); |
| auto currentAngle = angles[0].dyn_cast<FloatAttr>(); |
| if (!currentAngle) |
| return failure(); |
|
|
| |
| Value qubit = op.getQubits()[0]; |
| auto prevOp = qubit.getDefiningOp<UnitaryOp>(); |
| if (!prevOp || prevOp.getQubits().size() != 1) |
| return failure(); |
| if (prevOp.getAxis() && *prevOp.getAxis() != "Z") |
| return failure(); |
| auto prevAngles = prevOp.getAngles(); |
| if (prevAngles.size() != 1) |
| return failure(); |
| auto prevAngle = prevAngles[0].dyn_cast<FloatAttr>(); |
| if (!prevAngle) |
| return failure(); |
| if (prevOp.getQubits()[0] != qubit) |
| return failure(); |
|
|
| |
| double combined = currentAngle.getValueAsDouble() + |
| prevAngle.getValueAsDouble(); |
|
|
| |
| auto loc = op.getLoc(); |
| auto newAngle = rewriter.getFloatAttr(rewriter.getF64Type(), combined); |
| auto newAngles = rewriter.getArrayAttr({newAngle}); |
| auto zAxis = rewriter.getStringAttr("Z"); |
|
|
| auto combinedOp = rewriter.create<UnitaryOp>( |
| loc, TypeRange{qubit.getType()}, newAngles, zAxis, |
| ValueRange{qubit}); |
|
|
| rewriter.replaceOp(op, combinedOp.getResults()); |
| return success(); |
| } |
| }; |
|
|
| |
| |
| |
| struct DoubleZCancellation : public OpRewritePattern<UnitaryOp> { |
| using OpRewritePattern::OpRewritePattern; |
|
|
| LogicalResult matchAndRewrite(UnitaryOp op, |
| PatternRewriter &rewriter) const override { |
| if (op.getQubits().size() != 1) |
| return failure(); |
| auto angles = op.getAngles(); |
| if (angles.size() != 1) |
| return failure(); |
| auto angle = angles[0].dyn_cast<FloatAttr>(); |
| if (!angle) |
| return failure(); |
|
|
| |
| if (std::abs(angle.getValueAsDouble() - 0.5) > 1e-10) |
| return failure(); |
| if (!op.getAxis() || *op.getAxis() != "Z") |
| return failure(); |
|
|
| |
| Value qubit = op.getQubits()[0]; |
| auto prevOp = qubit.getDefiningOp<UnitaryOp>(); |
| if (!prevOp || prevOp.getQubits().size() != 1) |
| return failure(); |
| if (prevOp.getQubits()[0] != qubit) |
| return failure(); |
| auto prevAngles = prevOp.getAngles(); |
| if (prevAngles.size() != 1) |
| return failure(); |
| auto prevAngle = prevAngles[0].dyn_cast<FloatAttr>(); |
| if (!prevAngle) |
| return failure(); |
| if (std::abs(prevAngle.getValueAsDouble() - 0.5) > 1e-10) |
| return failure(); |
| if (!prevOp.getAxis() || *prevOp.getAxis() != "Z") |
| return failure(); |
|
|
| |
| rewriter.replaceOp(op, prevOp.getQubits()); |
| return success(); |
| } |
| }; |
|
|
| |
| |
| |
|
|
| void mlir::quantum::populateQuantumRewritePatterns( |
| mlir::RewritePatternSet &patterns, MLIRContext *ctx) { |
| patterns.add<HHCancellation>(ctx); |
| patterns.add<TripleTCancellation>(ctx); |
| patterns.add<IdentityElimination>(ctx); |
| patterns.add<RzCancellation>(ctx); |
| patterns.add<DoubleZCancellation>(ctx); |
| } |
|
|
| |
| |
| |
|
|
| LogicalResult mlir::quantum::applyQuantumRewrites(func::FuncOp funcOp) { |
| MLIRContext *ctx = funcOp.getContext(); |
| RewritePatternSet patterns(ctx); |
| populateQuantumRewritePatterns(patterns, ctx); |
|
|
| GreedyRewriteConfig config; |
| config.useTopDownTraversal = true; |
| config.maxIterations = 100; |
|
|
| return applyPatternsAndFoldGreedily(funcOp, std::move(patterns), config); |
| } |
|
|