diff --git a/bindings/mlir/qiskit/Qiskit2_5.cpp b/bindings/mlir/qiskit/Qiskit2_5.cpp index cecf5e9664..ccb66db6b5 100644 --- a/bindings/mlir/qiskit/Qiskit2_5.cpp +++ b/bindings/mlir/qiskit/Qiskit2_5.cpp @@ -1189,14 +1189,16 @@ class NativeControlFlowReader final : public ControlFlowReader { throw std::runtime_error( "Qiskit classical-bit condition must compare against zero or one"); } - result.expectedBit = expected != 0U; - return result; + return normalizePythonTarget(expressionModule.attr("equal")( + condition[0], nb::bool_(expected != 0U))); } if (result.kind == ClassicalTargetKind::ClassicalRegister) { - result.width = static_cast( - std::max(result.reg.bits.size(), std::bit_width(expected))); - result.expectedRegister = expected; - return result; + if (std::bit_width(expected) > result.reg.bits.size()) { + return normalizePythonTarget( + expressionModule.attr("lift")(nb::bool_(false))); + } + return normalizePythonTarget( + expressionModule.attr("equal")(condition[0], nb::int_(expected))); } throw std::runtime_error("Qiskit control flow has an unknown condition " "target"); @@ -1646,34 +1648,15 @@ class PythonClassicalBuilder final { } [[nodiscard]] nb::object condition(const ClassicalTarget& target) const { - switch (target.kind) { - case ClassicalTargetKind::ClassicalBit: - return nb::make_tuple(classicalBit(target.bit), - nb::bool_(target.expectedBit)); - case ClassicalTargetKind::ClassicalRegister: { - validateRegisterValue(target.reg, target.expectedRegister); - if (const auto reg = registeredClassicalRegister(target.reg)) { - return nb::make_tuple(*reg, nb::int_(target.expectedRegister)); - } - const auto packed = packedRegister(target.reg); - const auto expected = expressionModule_.attr("lift")( - nb::int_(target.expectedRegister), - classicalType(ClassicalType::Uint, - static_cast(target.reg.bits.size()))); - return expressionModule_.attr("equal")(packed, expected); + if (target.kind != ClassicalTargetKind::Expression || !target.expression) { + throw std::runtime_error( + "Qiskit control-flow condition has no expression"); } - case ClassicalTargetKind::Expression: - if (!target.expression) { - throw std::runtime_error( - "Qiskit control-flow condition has no expression"); - } - if (target.expression->type != ClassicalType::Bool) { - throw std::runtime_error( - "Qiskit control-flow condition expression must be Boolean"); - } - return expression(*target.expression); + if (target.expression->type != ClassicalType::Bool) { + throw std::runtime_error( + "Qiskit control-flow condition expression must be Boolean"); } - throw std::runtime_error("Qiskit control flow has an unknown condition"); + return expression(*target.expression); } [[nodiscard]] nb::object switchTarget(const ClassicalTarget& target) const { @@ -1752,18 +1735,6 @@ class PythonClassicalBuilder final { return std::nullopt; } - static void validateRegisterValue(const Register& reg, const uint64_t value) { - if (reg.bits.empty() || reg.bits.size() > 64U) { - throw std::runtime_error( - "Qiskit condition registers must contain between 1 and 64 bits"); - } - if (reg.bits.size() < std::numeric_limits::digits && - value >= (uint64_t{1} << reg.bits.size())) { - throw std::runtime_error( - "Qiskit register condition value exceeds its register width"); - } - } - [[nodiscard]] nb::object packedRegister(const Register& reg, const uint32_t expressionWidth = 0U) const { diff --git a/bindings/mlir/qiskit/QiskitExport.cpp b/bindings/mlir/qiskit/QiskitExport.cpp index 4c2aacd1d2..da5280c355 100644 --- a/bindings/mlir/qiskit/QiskitExport.cpp +++ b/bindings/mlir/qiskit/QiskitExport.cpp @@ -977,6 +977,53 @@ static void setExpressionType(Expression& expression, const mlir::Type type) { return checkedAdd(info->second.base, checked, "classical-bit"); } +[[nodiscard]] static Register classicalRegister(mlir::Value value, + const ExportState& state) { + const auto info = state.classicalRegisterInfo.find(value); + if (info == state.classicalRegisterInfo.end() || info->second.size == 0U || + info->second.size > 64U) { + throw std::runtime_error( + "Qiskit register comparisons require between 1 and 64 bits"); + } + if (info->second.initialization != mlir::cbit::Initialization::Zero) { + const auto written = state.unconditionalWrites.find(value); + if (written == state.unconditionalWrites.end() || + written->second.size() != info->second.size) { + throw std::runtime_error( + "Qiskit register comparison reads undefined classical bits"); + } + } + Register result; + result.bits.resize(info->second.size); + std::iota(result.bits.begin(), result.bits.end(), info->second.base); + for (const auto& candidate : state.classicalRegisters) { + if (candidate.bits == result.bits) { + result.name = candidate.name; + break; + } + } + return result; +} + +[[nodiscard]] static BinaryOperation +comparisonOperation(const mlir::cbit::ComparisonPredicate predicate) { + switch (predicate) { + case mlir::cbit::ComparisonPredicate::Equal: + return BinaryOperation::Equal; + case mlir::cbit::ComparisonPredicate::NotEqual: + return BinaryOperation::NotEqual; + case mlir::cbit::ComparisonPredicate::Less: + return BinaryOperation::Less; + case mlir::cbit::ComparisonPredicate::LessEqual: + return BinaryOperation::LessEqual; + case mlir::cbit::ComparisonPredicate::Greater: + return BinaryOperation::Greater; + case mlir::cbit::ComparisonPredicate::GreaterEqual: + return BinaryOperation::GreaterEqual; + } + llvm_unreachable("unknown CBit comparison predicate"); +} + [[noreturn]] static void throwClassicalExpressionSizeError() { throw std::runtime_error( "QC classical expression exceeds the size limit of 4096 nodes"); @@ -1073,6 +1120,31 @@ exportExpressionImpl(mlir::Value value, ExportState& state, state.expressionOperations.insert(operation); return result; } + if (auto comparison = llvm::dyn_cast(operation)) { + const auto width = comparison.getRhs().getBitWidth(); + if (width > 64U) { + throw std::runtime_error( + "Qiskit register comparisons support at most 64 bits"); + } + countExpressionNode(nodeCount); + auto left = std::make_unique(); + left->kind = ExpressionKind::ClassicalRegister; + left->type = ClassicalType::Uint; + left->width = width; + left->reg = classicalRegister(comparison.getReg(), state); + countExpressionNode(nodeCount); + auto right = std::make_unique(); + right->kind = ExpressionKind::Value; + right->type = ClassicalType::Uint; + right->width = width; + right->uintValue = comparison.getRhs().getZExtValue(); + result->kind = ExpressionKind::Binary; + result->binaryOperation = comparisonOperation(comparison.getPredicate()); + result->left = std::move(left); + result->right = std::move(right); + state.expressionOperations.insert(operation); + return result; + } if (auto ifOp = llvm::dyn_cast(operation)) { if (ifOp.getNumResults() != 1U || !value.getType().isInteger(1) || ifOp.getElseRegion().empty()) { @@ -1393,7 +1465,7 @@ static void acceptPackedRegister(PackedRegister& packed, ExportState& state) { static void validateClassicalSnapshot(mlir::Value expression, mlir::Operation& consumer) { llvm::DenseSet visited; - llvm::SmallVector loads; + llvm::SmallVector> reads; llvm::SmallVector worklist{expression}; while (!worklist.empty()) { auto value = worklist.pop_back_val(); @@ -1408,7 +1480,11 @@ static void validateClassicalSnapshot(mlir::Value expression, continue; } if (auto load = llvm::dyn_cast(operation)) { - loads.push_back(load); + reads.emplace_back(load, load.getReg()); + continue; + } + if (auto comparison = llvm::dyn_cast(operation)) { + reads.emplace_back(comparison, comparison.getReg()); continue; } if (auto ifOp = llvm::dyn_cast(operation)) { @@ -1422,9 +1498,9 @@ static void validateClassicalSnapshot(mlir::Value expression, } worklist.append(operation->operand_begin(), operation->operand_end()); } - for (auto load : loads) { - mlir::Operation* anchor = load; - auto* anchorBlock = load->getBlock(); + for (auto [read, reg] : reads) { + mlir::Operation* anchor = read; + auto* anchorBlock = read->getBlock(); while (anchorBlock != consumer.getBlock()) { auto* parent = anchorBlock->getParentOp(); auto parentIf = llvm::dyn_cast_if_present(parent); @@ -1448,13 +1524,13 @@ static void validateClassicalSnapshot(mlir::Value expression, "Qiskit control-flow expression does not dominate its consumer"); } if (auto store = llvm::dyn_cast(operation); - store && store.getReg() == load.getReg()) { + store && store.getReg() == reg) { throw std::runtime_error( "Qiskit control-flow export cannot preserve a stale classical " "snapshot"); } if (operation->getNumRegions() != 0U && - storesToValueRecursively(*operation, load.getReg())) { + storesToValueRecursively(*operation, reg)) { throw std::runtime_error( "Qiskit control-flow export cannot preserve a classical " "snapshot across nested control flow"); @@ -1471,39 +1547,6 @@ exportCondition(mlir::Value value, ExportState& state, "Qiskit control-flow conditions must have Boolean type"); } validateClassicalSnapshot(value, consumer); - if (auto comparison = value.getDefiningOp(); - comparison && - comparison.getPredicate() == mlir::arith::CmpIPredicate::eq) { - for (auto [actual, expected] : - std::array{std::pair{comparison.getLhs(), comparison.getRhs()}, - std::pair{comparison.getRhs(), comparison.getLhs()}}) { - const auto constant = constantUnsignedInteger(expected); - if (!constant) { - continue; - } - if (auto load = actual.getDefiningOp(); - load && actual.getType().isInteger(1) && *constant <= 1U) { - state.expressionOperations.insert(comparison); - state.expressionOperations.insert(load); - return {.kind = ClassicalTargetKind::ClassicalBit, - .bit = classicalBitIndex(load, state), - .expectedBit = *constant != 0U}; - } - if (auto packed = matchPackedRegister(actual, state, evaluationBlock)) { - if (packed->reg.bits.size() != 64U && - *constant >= (uint64_t{1} << packed->reg.bits.size())) { - continue; - } - state.expressionOperations.insert(comparison); - acceptPackedRegister(*packed, state); - return {.kind = ClassicalTargetKind::ClassicalRegister, - .reg = std::move(packed->reg), - .expectedRegister = *constant, - .width = - llvm::cast(actual.getType()).getWidth()}; - } - } - } ClassicalTarget target{.kind = ClassicalTargetKind::Expression}; target.expression = exportExpression(value, state, evaluationBlock); return target; @@ -1866,6 +1909,10 @@ collectSwitch(mlir::scf::IndexSwitchOp switchOp, ExportState& state, deferredExpressions.push_back(&operation); continue; } + if (llvm::isa(operation)) { + deferredExpressions.push_back(&operation); + continue; + } if (auto dealloc = llvm::dyn_cast(operation)) { if (topLevel && state.quantumBases.contains(dealloc.getMemref())) { continue; diff --git a/bindings/mlir/qiskit/QiskitImport.cpp b/bindings/mlir/qiskit/QiskitImport.cpp index 2fe44d3594..40d0a8d61b 100644 --- a/bindings/mlir/qiskit/QiskitImport.cpp +++ b/bindings/mlir/qiskit/QiskitImport.cpp @@ -15,6 +15,7 @@ #include "jeff/IR/JeffDialect.h" #include "mlir/Compiler/Programs.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" +#include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/MQT/IR/MQTDialect.h" #include "mlir/Dialect/MQT/Utils/DenseUnitary.h" #include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" @@ -620,20 +621,55 @@ struct ClassicalBitRef { }; } // namespace -[[nodiscard]] static mlir::Value -loadClassicalBit(mlir::qc::QCProgramBuilder& builder, - const llvm::ArrayRef classicalBits, - const llvm::ArrayRef rootClbitMap, - const uint32_t index) { +[[nodiscard]] static const ClassicalBitRef& +classicalBitRef(const llvm::ArrayRef classicalBits, + const llvm::ArrayRef rootClbitMap, uint32_t index) { if (index >= rootClbitMap.size() || rootClbitMap[index] >= classicalBits.size()) { throw std::runtime_error( "Qiskit control flow references an invalid classical bit"); } - const auto& bit = classicalBits[rootClbitMap[index]]; + return classicalBits[rootClbitMap[index]]; +} + +[[nodiscard]] static mlir::Value +loadClassicalBit(mlir::qc::QCProgramBuilder& builder, + const llvm::ArrayRef classicalBits, + const llvm::ArrayRef rootClbitMap, + const uint32_t index) { + const auto& bit = classicalBitRef(classicalBits, rootClbitMap, index); return builder.loadClassicalBit(bit.storage, bit.index); } +[[nodiscard]] static mlir::Value emitRegisterComparison( + mlir::qc::QCProgramBuilder& builder, + const llvm::ArrayRef classicalBits, + const llvm::ArrayRef rootClbitMap, const Register& reg, + mlir::cbit::ComparisonPredicate predicate, uint64_t expected) { + mlir::Value storage; + for (size_t index = 0U; index < reg.bits.size(); ++index) { + const auto& bit = + classicalBitRef(classicalBits, rootClbitMap, reg.bits[index]); + if (std::cmp_not_equal(bit.index, index) || + (storage && bit.storage != storage)) { + return {}; + } + storage = bit.storage; + } + if (!storage) { + return {}; + } + const auto width = static_cast(reg.bits.size()); + if (llvm::cast(storage.getType()).getWidth() != + width) { + return {}; + } + const auto rhs = builder.getIntegerAttr(builder.getIntegerType(width), + llvm::APInt(width, expected, false)); + return mlir::cbit::CompareOp::create(builder, builder.getI1Type(), predicate, + storage, rhs); +} + [[nodiscard]] static mlir::Value packRegister(mlir::qc::QCProgramBuilder& builder, const llvm::ArrayRef classicalBits, @@ -675,6 +711,31 @@ packRegister(mlir::qc::QCProgramBuilder& builder, return terms.front(); } +[[nodiscard]] static std::optional +registerComparisonPredicate(const BinaryOperation operation, + const bool reverse) { + switch (operation) { + case BinaryOperation::Equal: + return mlir::cbit::ComparisonPredicate::Equal; + case BinaryOperation::NotEqual: + return mlir::cbit::ComparisonPredicate::NotEqual; + case BinaryOperation::Less: + return reverse ? mlir::cbit::ComparisonPredicate::Greater + : mlir::cbit::ComparisonPredicate::Less; + case BinaryOperation::LessEqual: + return reverse ? mlir::cbit::ComparisonPredicate::GreaterEqual + : mlir::cbit::ComparisonPredicate::LessEqual; + case BinaryOperation::Greater: + return reverse ? mlir::cbit::ComparisonPredicate::Less + : mlir::cbit::ComparisonPredicate::Greater; + case BinaryOperation::GreaterEqual: + return reverse ? mlir::cbit::ComparisonPredicate::LessEqual + : mlir::cbit::ComparisonPredicate::GreaterEqual; + default: + return std::nullopt; + } +} + [[nodiscard]] static mlir::Value emitExpression(mlir::qc::QCProgramBuilder& builder, const Expression& expression, @@ -708,6 +769,15 @@ emitExpression(mlir::qc::QCProgramBuilder& builder, target); } case ExpressionKind::Cast: { + if (expression.type == ClassicalType::Bool && + expression.left->kind == ExpressionKind::ClassicalRegister && + expression.left->width == expression.left->reg.bits.size()) { + if (auto comparison = emitRegisterComparison( + builder, classicalBits, rootClbitMap, expression.left->reg, + mlir::cbit::ComparisonPredicate::NotEqual, 0U)) { + return comparison; + } + } auto operand = emitExpression(builder, *expression.left, classicalBits, rootClbitMap); if (operand.getType() == resultType) { @@ -784,6 +854,27 @@ emitExpression(mlir::qc::QCProgramBuilder& builder, break; } case ExpressionKind::Binary: { + const auto reverse = + expression.left->kind == ExpressionKind::Value && + expression.right->kind == ExpressionKind::ClassicalRegister; + const auto& registerExpression = + reverse ? *expression.right : *expression.left; + const auto& expected = reverse ? *expression.left : *expression.right; + if (const auto predicate = + registerComparisonPredicate(expression.binaryOperation, reverse); + predicate && + registerExpression.kind == ExpressionKind::ClassicalRegister && + registerExpression.type == ClassicalType::Uint && + registerExpression.width == registerExpression.reg.bits.size() && + expected.kind == ExpressionKind::Value && + expected.type == ClassicalType::Uint && + expected.width == registerExpression.width) { + if (auto comparison = emitRegisterComparison( + builder, classicalBits, rootClbitMap, registerExpression.reg, + *predicate, expected.uintValue)) { + return comparison; + } + } auto left = emitExpression(builder, *expression.left, classicalBits, rootClbitMap); if (expression.binaryOperation == BinaryOperation::LogicAnd || @@ -956,36 +1047,17 @@ emitCondition(mlir::qc::QCProgramBuilder& builder, const ClassicalTarget& target, const llvm::ArrayRef classicalBits, const llvm::ArrayRef rootClbitMap) { - switch (target.kind) { - case ClassicalTargetKind::ClassicalBit: { - auto actual = - loadClassicalBit(builder, classicalBits, rootClbitMap, target.bit); - return mlir::arith::CmpIOp::create(builder, mlir::arith::CmpIPredicate::eq, - actual, - builder.boolConstant(target.expectedBit)) - .getResult(); - } - case ClassicalTargetKind::ClassicalRegister: { - auto actual = castInteger( - builder, packRegister(builder, classicalBits, rootClbitMap, target.reg), - builder.getIntegerType(target.width)); - auto expected = - integerConstant(builder, target.width, target.expectedRegister); - return mlir::arith::CmpIOp::create(builder, mlir::arith::CmpIPredicate::eq, - actual, expected) - .getResult(); - } - case ClassicalTargetKind::Expression: { - auto condition = emitExpression(builder, *target.expression, classicalBits, - rootClbitMap); - if (!condition.getType().isInteger(1)) { - throw std::runtime_error( - "Qiskit control-flow condition expression must have Boolean type"); - } - return condition; + if (target.kind != ClassicalTargetKind::Expression || !target.expression) { + throw std::runtime_error( + "Qiskit control-flow condition has no classical expression"); } + auto condition = + emitExpression(builder, *target.expression, classicalBits, rootClbitMap); + if (!condition.getType().isInteger(1)) { + throw std::runtime_error( + "Qiskit control-flow condition expression must have Boolean type"); } - throw std::runtime_error("unknown normalized Qiskit condition type"); + return condition; } [[nodiscard]] static mlir::Value diff --git a/bindings/mlir/qiskit/QiskitTranslation.h b/bindings/mlir/qiskit/QiskitTranslation.h index a66a20e5d3..fa5e442073 100644 --- a/bindings/mlir/qiskit/QiskitTranslation.h +++ b/bindings/mlir/qiskit/QiskitTranslation.h @@ -277,9 +277,7 @@ enum class ClassicalTargetKind : uint8_t { struct ClassicalTarget { ClassicalTargetKind kind = ClassicalTargetKind::ClassicalBit; uint32_t bit = 0; - bool expectedBit = false; Register reg; - uint64_t expectedRegister = 0; uint32_t width = 1; std::unique_ptr expression; }; diff --git a/docs/mlir/OpenQASM.md b/docs/mlir/OpenQASM.md index a8244c56e1..62bc657e93 100644 --- a/docs/mlir/OpenQASM.md +++ b/docs/mlir/OpenQASM.md @@ -52,6 +52,10 @@ Runtime integer preconditions and classical-index bounds are represented explicitly in QC. This safety machinery is supported by the normal compiler and QIR paths, but it is intentionally outside the export subset described below. +OpenQASM 3 supports all six unsigned comparisons with a bit register on the left +and an integer constant on the right. OpenQASM 2 retains its equality-only +register condition. + Fixed-width angles are a compile-time input feature. An omitted angle width resolves to 52 bits. Both `const angle[N]` and initialized `angle[N]` declarations are accepted as write-once values. Initializers and angle casts @@ -83,10 +87,12 @@ do not index qubits keep their runtime behavior. Bit registers use `!cbit.reg` in QC. OpenQASM 2 initializes each register to zero. OpenQASM 3 leaves each register undefined until a statement writes it. -Explicit outputs and implicit global outputs are returned by the entry function; -internal CBit allocations are not outputs. Other scalar outputs use builtin MLIR -scalar types. A scalar `qubit` lowers to `qc.alloc`, while `qubit[1]` remains a -one-element qubit register. +Whole-register comparisons lower to `cbit.cmp` and keep their unsigned integer +meaning without expanding into per-bit expression trees. Explicit outputs and +implicit global outputs are returned by the entry function; internal CBit +allocations are not outputs. Other scalar outputs use builtin MLIR scalar types. +A scalar `qubit` lowers to `qc.alloc`, while `qubit[1]` remains a one-element +qubit register. ## Export OpenQASM @@ -187,6 +193,7 @@ Unsigned constants therefore normalize to `int`. Operations whose signedness affects their meaning, such as unsigned division, comparison, or conversion, are rejected instead of being approximated. Integer sign extension and truncation are also rejected because OpenQASM scalar casts have different value semantics. +Direct `cbit.cmp` operations retain their unsigned register semantics. Emitted scalar casts use standard OpenQASM conversion syntax. The MQT Core frontend does not yet parse that syntax, so cast-containing output is outside diff --git a/mlir/include/mlir/Conversion/QCToQIR/QIRCommon/QIRCommon.h b/mlir/include/mlir/Conversion/QCToQIR/QIRCommon/QIRCommon.h index 34fb68d5f1..63dcda4035 100644 --- a/mlir/include/mlir/Conversion/QCToQIR/QIRCommon/QIRCommon.h +++ b/mlir/include/mlir/Conversion/QCToQIR/QIRCommon/QIRCommon.h @@ -40,6 +40,9 @@ struct LoweringState { /// Result-array pointers to be deallocated at the end of the program DenseSet resultArrays; + /// CBit read operations whose register is backed by a result array. + DenseSet returnedCBitReads; + /// Cache static qubit pointers for reuse DenseMap staticQubits; diff --git a/mlir/include/mlir/Dialect/CBit/IR/CBitOps.h b/mlir/include/mlir/Dialect/CBit/IR/CBitOps.h index a4e0bfe508..c0686ad430 100644 --- a/mlir/include/mlir/Dialect/CBit/IR/CBitOps.h +++ b/mlir/include/mlir/Dialect/CBit/IR/CBitOps.h @@ -14,6 +14,7 @@ #include "mlir/Dialect/CBit/IR/CBitDialect.h" #include +#include #include #include @@ -29,4 +30,9 @@ namespace mlir::cbit { void validateStaticRegisterIndex(Value reg, const std::variant& index); +/// Builds an equivalent comparison from individual register bits. +Value buildComparison(OpBuilder& builder, Location location, + ComparisonPredicate predicate, const llvm::APInt& rhs, + llvm::function_ref loadBit); + } // namespace mlir::cbit diff --git a/mlir/include/mlir/Dialect/CBit/IR/CBitOps.td b/mlir/include/mlir/Dialect/CBit/IR/CBitOps.td index 3d0450f78a..3b530490a7 100644 --- a/mlir/include/mlir/Dialect/CBit/IR/CBitOps.td +++ b/mlir/include/mlir/Dialect/CBit/IR/CBitOps.td @@ -33,6 +33,17 @@ def CBit_InitializationAttr let assemblyFormat = "`<` $value `>`"; } +def CBit_ComparisonPredicate + : I64EnumAttr<"ComparisonPredicate", "unsigned comparison predicate", + [I64EnumAttrCase<"Equal", 0, "eq">, + I64EnumAttrCase<"NotEqual", 1, "ne">, + I64EnumAttrCase<"Less", 2, "ult">, + I64EnumAttrCase<"LessEqual", 3, "ule">, + I64EnumAttrCase<"Greater", 4, "ugt">, + I64EnumAttrCase<"GreaterEqual", 5, "uge">]> { + let cppNamespace = "::mlir::cbit"; +} + def CBit_RegisterType : TypeDef { let mnemonic = "reg"; let summary = "A static classical-bit register"; @@ -96,6 +107,28 @@ def LoadOp : CBitOp<"load"> { let hasVerifier = 1; } +def CompareOp : CBitOp<"cmp"> { + let summary = "Compare a classical-bit register with an integer"; + let description = [{ + Compares the complete register with an unsigned integer of the same width. + The predicate must be `eq`, `ne`, `ult`, `ule`, `ugt`, or `uge`. + + Example: + ```mlir + %matches = cbit.cmp eq, %c, 1 : i2 : !cbit.reg<2> + ``` + }]; + + let arguments = (ins CBit_ComparisonPredicate:$predicate, + Arg:$reg, APIntAttr:$rhs); + let results = (outs I1:$result); + let assemblyFormat = [{ + $predicate `,` $reg `,` $rhs attr-dict `:` qualified(type($reg)) + }]; + let hasCanonicalizer = 1; + let hasVerifier = 1; +} + def StoreOp : CBitOp<"store"> { let summary = "Store a classical bit"; let description = [{ diff --git a/mlir/include/mlir/Target/OpenQASM/Frontend.h b/mlir/include/mlir/Target/OpenQASM/Frontend.h index 4bedb090a5..6b071cadb1 100644 --- a/mlir/include/mlir/Target/OpenQASM/Frontend.h +++ b/mlir/include/mlir/Target/OpenQASM/Frontend.h @@ -10,6 +10,7 @@ #pragma once +#include #include #include @@ -208,6 +209,7 @@ enum class ConditionKind : uint8_t { Not, And, Or, + RegisterComparison, Comparison, }; @@ -220,6 +222,8 @@ struct ConditionExpression { QubitReference measurement; ConditionId lhs = 0; ConditionId rhs = 0; + RegisterId reg = 0; + llvm::APInt expected = llvm::APInt(1, 0); ExpressionId comparisonLhs = 0; ExpressionId comparisonRhs = 0; ComparisonKind comparison = ComparisonKind::Equal; diff --git a/mlir/lib/Conversion/CBitToMemRef/CBitToMemRef.cpp b/mlir/lib/Conversion/CBitToMemRef/CBitToMemRef.cpp index dd4ffa5982..3840441811 100644 --- a/mlir/lib/Conversion/CBitToMemRef/CBitToMemRef.cpp +++ b/mlir/lib/Conversion/CBitToMemRef/CBitToMemRef.cpp @@ -88,6 +88,25 @@ struct ConvertLoadOp final : OpConversionPattern { } }; +struct ConvertCompareOp final : OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(cbit::CompareOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + auto result = cbit::buildComparison( + rewriter, op.getLoc(), op.getPredicate(), op.getRhs(), + [&](const int64_t index) -> Value { + auto indexValue = + arith::ConstantIndexOp::create(rewriter, op.getLoc(), index); + return memref::LoadOp::create(rewriter, op.getLoc(), adaptor.getReg(), + ValueRange{indexValue}); + }); + rewriter.replaceOp(op, result); + return success(); + } +}; + struct ConvertStoreOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -123,8 +142,9 @@ struct ConvertCBitToMemRef final target.addDynamicallyLegalOp( [&](Operation* op) { return typeConverter.isLegal(op); }); - patterns.add(typeConverter, - context); + patterns + .add( + typeConverter, context); populateFunctionOpInterfaceTypeConversionPattern( patterns, typeConverter); populateReturnOpTypeConversionPattern(patterns, typeConverter); diff --git a/mlir/lib/Conversion/QCOToJeff/QCOToJeff.cpp b/mlir/lib/Conversion/QCOToJeff/QCOToJeff.cpp index 505c54cd5c..97a77e0aab 100644 --- a/mlir/lib/Conversion/QCOToJeff/QCOToJeff.cpp +++ b/mlir/lib/Conversion/QCOToJeff/QCOToJeff.cpp @@ -631,6 +631,35 @@ struct ConvertCBitLoadOpToJeff final } }; +/// Converts a CBit register comparison to jeff array reads and Boolean logic. +struct ConvertCBitCompareOpToJeff final + : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(cbit::CompareOp op, OpAdaptor /*adaptor*/, + ConversionPatternRewriter& rewriter) const override { + auto& state = getState().cbitState; + auto reg = state.resolveRegisterUse(op, op.getReg()); + auto array = state.getCurrentValue(reg, op); + if (!array) { + return rewriter.notifyMatchFailure(op, "unknown classical register"); + } + array = rewriter.getRemappedValue(array); + auto result = cbit::buildComparison( + rewriter, op.getLoc(), op.getPredicate(), op.getRhs(), + [&](const int64_t index) -> Value { + auto position = jeff::IntConst32Op::create( + rewriter, op.getLoc(), + rewriter.getI32IntegerAttr(static_cast(index))); + return jeff::IntArrayGetIndexOp::create( + rewriter, op.getLoc(), rewriter.getI1Type(), array, position); + }); + rewriter.replaceOp(op, result); + return success(); + } +}; + /** * @brief Converts qtensor.alloc to jeff.qureg_alloc * @@ -1879,13 +1908,14 @@ struct QCOToJeff final : impl::QCOToJeffBase { // Register operation conversion patterns jeff::populateNativeToJeffConversionPatterns(patterns); - patterns.add(typeConverter, context, &state); + patterns.add( + typeConverter, context, &state); using JK = JeffKind; using PP = PPRPaulis; diff --git a/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp b/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp index f00d37b32c..0bd7f02c9c 100644 --- a/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp +++ b/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp @@ -649,9 +649,9 @@ collectRegisterAccesses(Operation* root, LoweringState& state) { } } - if (!isa(operation)) { + if (!isa(operation)) { return WalkResult::advance(); } diff --git a/mlir/lib/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.cpp b/mlir/lib/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.cpp index 809d724706..f234c04b64 100644 --- a/mlir/lib/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.cpp +++ b/mlir/lib/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.cpp @@ -40,6 +40,7 @@ #include #include #include +#include #include #include #include @@ -56,6 +57,109 @@ using namespace qir; #define GEN_PASS_DEF_QCTOQIRADAPTIVE #include "mlir/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.h.inc" +namespace { + +constexpr unsigned LOCAL_CBIT_REGISTER = 1U; +constexpr unsigned RETURNED_CBIT_REGISTER = 2U; +constexpr unsigned MIXED_CBIT_REGISTER = + LOCAL_CBIT_REGISTER | RETURNED_CBIT_REGISTER; + +} // namespace + +static LogicalResult prepareCBitRegisterReads(Operation* moduleOp, + LoweringState& state) { + DenseMap> forwardedRegisters; + moduleOp->walk([&](Operation* operation) { + for (auto& region : operation->getRegions()) { + for (auto& block : region) { + for (auto argument : block.getArguments()) { + if (!isa(argument.getType())) { + continue; + } + for (auto* predecessor : block.getPredecessors()) { + auto branch = + dyn_cast(predecessor->getTerminator()); + if (!branch) { + continue; + } + for (unsigned successorIndex = 0; + successorIndex < branch->getNumSuccessors(); + ++successorIndex) { + if (branch->getSuccessor(successorIndex) != &block) { + continue; + } + auto operands = branch.getSuccessorOperands(successorIndex); + if (argument.getArgNumber() >= operands.size() || + operands.isOperandProduced(argument.getArgNumber())) { + continue; + } + if (auto incoming = operands[argument.getArgNumber()]) { + forwardedRegisters[incoming].push_back(argument); + } + } + } + } + } + } + }); + moduleOp->walk([&](arith::SelectOp selectOp) { + if (isa(selectOp.getType())) { + forwardedRegisters[selectOp.getTrueValue()].push_back( + selectOp.getResult()); + forwardedRegisters[selectOp.getFalseValue()].push_back( + selectOp.getResult()); + } + }); + + DenseMap representations; + SmallVector worklist; + moduleOp->walk([&](cbit::AllocOp allocOp) { + const auto it = state.cregIndices.find(allocOp.getOperation()); + if (it == state.cregIndices.end()) { + return; + } + representations[allocOp.getResult()] = state.cregs[it->second].record + ? RETURNED_CBIT_REGISTER + : LOCAL_CBIT_REGISTER; + worklist.push_back(allocOp.getResult()); + }); + + while (!worklist.empty()) { + auto source = worklist.pop_back_val(); + const auto it = forwardedRegisters.find(source); + if (it == forwardedRegisters.end()) { + continue; + } + for (auto destination : it->second) { + auto& representation = representations[destination]; + const auto merged = representation | representations.lookup(source); + if (merged != representation) { + representation = merged; + worklist.push_back(destination); + } + } + } + + bool hasMixedRepresentation = false; + const auto prepareRead = [&](Operation* operation, Value reg) { + const auto representation = representations.lookup(reg); + if (representation == MIXED_CBIT_REGISTER) { + operation->emitOpError( + "adaptive QIR conversion cannot merge returned and local CBit " + "registers"); + hasMixedRepresentation = true; + } else if (representation == RETURNED_CBIT_REGISTER) { + state.returnedCBitReads.insert(operation); + } + }; + moduleOp->walk( + [&](cbit::LoadOp loadOp) { prepareRead(loadOp, loadOp.getReg()); }); + moduleOp->walk([&](cbit::CompareOp compareOp) { + prepareRead(compareOp, compareOp.getReg()); + }); + return success(!hasMixedRepresentation); +} + /** * @brief Returns the result pointer the `qc::MeasureOp` @p op writes to, or * `nullptr` if it does not write into a classical register. @@ -187,33 +291,64 @@ struct ConvertCBitAllocOp final : StatefulOpConversionPattern { } }; +} // namespace + +static Value loadCBit(Operation* op, Value reg, Value index, + ConversionPatternRewriter& rewriter, + bool returnedRegister) { + const auto ptrType = LLVM::LLVMPointerType::get(rewriter.getContext()); + if (!returnedRegister) { + auto elementptr = + LLVM::GEPOp::create(rewriter, op->getLoc(), ptrType, + rewriter.getI1Type(), reg, ValueRange{index}); + return LLVM::LoadOp::create(rewriter, op->getLoc(), rewriter.getI1Type(), + elementptr); + } + auto elementptr = LLVM::GEPOp::create(rewriter, op->getLoc(), ptrType, + ptrType, reg, ValueRange{index}); + auto result = + LLVM::LoadOp::create(rewriter, op->getLoc(), ptrType, elementptr); + auto fnSig = LLVM::LLVMFunctionType::get(rewriter.getI1Type(), {ptrType}); + auto fnDec = + getOrCreateFunctionDeclaration(rewriter, op, QIR_READ_RESULT, fnSig); + return LLVM::CallOp::create(rewriter, op->getLoc(), fnDec, result.getResult()) + .getResult(); +} + +namespace { + struct ConvertCBitLoadOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; LogicalResult matchAndRewrite(cbit::LoadOp op, OpAdaptor adaptor, ConversionPatternRewriter& rewriter) const override { - auto& state = getState(); - const auto ptrType = LLVM::LLVMPointerType::get(getContext()); - if (!state.resultArrays.contains(adaptor.getReg())) { - auto elementptr = LLVM::GEPOp::create( - rewriter, op.getLoc(), ptrType, rewriter.getI1Type(), - adaptor.getReg(), ValueRange{adaptor.getIndex()}); - rewriter.replaceOpWithNewOp(op, rewriter.getI1Type(), - elementptr); - return success(); - } - auto elementptr = - LLVM::GEPOp::create(rewriter, op.getLoc(), ptrType, ptrType, - adaptor.getReg(), ValueRange{adaptor.getIndex()}); - auto result = - LLVM::LoadOp::create(rewriter, op.getLoc(), ptrType, elementptr); - auto fnSig = LLVM::LLVMFunctionType::get(rewriter.getI1Type(), {ptrType}); - auto fnDec = - getOrCreateFunctionDeclaration(rewriter, op, QIR_READ_RESULT, fnSig); - auto readResult = - LLVM::CallOp::create(rewriter, op.getLoc(), fnDec, result.getResult()); - rewriter.replaceOp(op, readResult.getResult()); + const auto returnedRegister = + getState().returnedCBitReads.contains(op.getOperation()); + rewriter.replaceOp(op, loadCBit(op, adaptor.getReg(), adaptor.getIndex(), + rewriter, returnedRegister)); + return success(); + } +}; + +struct ConvertCBitCompareOp final + : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(cbit::CompareOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + const auto returnedRegister = + getState().returnedCBitReads.contains(op.getOperation()); + auto result = cbit::buildComparison( + rewriter, op.getLoc(), op.getPredicate(), op.getRhs(), + [&](const int64_t index) -> Value { + auto indexValue = LLVM::ConstantOp::create( + rewriter, op.getLoc(), rewriter.getI64Type(), index); + return loadCBit(op, adaptor.getReg(), indexValue, rewriter, + returnedRegister); + }); + rewriter.replaceOp(op, result); return success(); } }; @@ -542,8 +677,8 @@ static void populateQCToQIRAdaptivePatterns(RewritePatternSet& patterns, MLIRContext* ctx, LoweringState& state) { populateQCToQIRPatterns(patterns, typeConverter, ctx, state); - patterns.add(typeConverter, ctx, &state); @@ -723,6 +858,10 @@ struct QCToQIRAdaptive final : impl::QCToQIRAdaptiveBase { signalPassFailure(); return; } + if (failed(prepareCBitRegisterReads(moduleOp, state))) { + signalPassFailure(); + return; + } // Stage 2.1: Convert func dialect to LLVM { diff --git a/mlir/lib/Conversion/QCToQIR/QIRBase/QCToQIRBase.cpp b/mlir/lib/Conversion/QCToQIR/QIRBase/QCToQIRBase.cpp index a7f07325d0..e0ce975244 100644 --- a/mlir/lib/Conversion/QCToQIR/QIRBase/QCToQIRBase.cpp +++ b/mlir/lib/Conversion/QCToQIR/QIRBase/QCToQIRBase.cpp @@ -142,6 +142,17 @@ struct RejectCBitLoadOp final : OpConversionPattern { } }; +struct RejectCBitCompareOp final : OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(cbit::CompareOp op, OpAdaptor /*adaptor*/, + ConversionPatternRewriter& /*rewriter*/) const override { + return op.emitError( + "QIR Base Profile does not support classical-register comparisons"); + } +}; + struct ConvertMemRefAllocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -335,7 +346,7 @@ static void populateQCToQIRBasePatterns(RewritePatternSet& patterns, patterns.add(typeConverter, ctx, &state); - patterns.add(typeConverter, ctx); + patterns.add(typeConverter, ctx); } namespace { diff --git a/mlir/lib/Dialect/CBit/IR/CBitOps.cpp b/mlir/lib/Dialect/CBit/IR/CBitOps.cpp index 2268c55d2a..f25a82b7c7 100644 --- a/mlir/lib/Dialect/CBit/IR/CBitOps.cpp +++ b/mlir/lib/Dialect/CBit/IR/CBitOps.cpp @@ -18,6 +18,7 @@ #include #include #include +#include #include #include // IWYU pragma: keep #include @@ -26,6 +27,7 @@ #include #include +#include #include using namespace mlir; @@ -168,17 +170,115 @@ struct ForwardKnownLoad final : OpRewritePattern { return success(); } }; + +struct FoldUntouchedZeroComparison final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(CompareOp compare, + PatternRewriter& rewriter) const override { + auto alloc = compare.getReg().getDefiningOp(); + if (!alloc || alloc.getInitialization() != Initialization::Zero || + alloc->getBlock() != compare->getBlock() || + !alloc->isBeforeInBlock(compare)) { + return failure(); + } + for (auto* user : compare.getReg().getUsers()) { + if (!isa(user)) { + auto* ancestor = compare->getBlock()->findAncestorOpInBlock(*user); + if (ancestor != nullptr && ancestor->isBeforeInBlock(compare)) { + return failure(); + } + } + } + const auto zero = compare.getRhs().isZero(); + const auto result = [&] { + switch (compare.getPredicate()) { + case ComparisonPredicate::Equal: + case ComparisonPredicate::GreaterEqual: + return zero; + case ComparisonPredicate::NotEqual: + case ComparisonPredicate::Less: + return !zero; + case ComparisonPredicate::LessEqual: + return true; + case ComparisonPredicate::Greater: + return false; + } + llvm_unreachable("unknown CBit comparison predicate"); + }(); + rewriter.replaceOpWithNewOp(compare, result, 1); + return success(); + } +}; + } // namespace LogicalResult LoadOp::verify() { return verifyIndex(getOperation(), getReg(), getIndex()); } +LogicalResult CompareOp::verify() { + if (std::cmp_not_equal(getRhs().getBitWidth(), + getReg().getType().getWidth())) { + return emitOpError("expected integer width must match register width"); + } + return success(); +} + +Value mlir::cbit::buildComparison( + OpBuilder& builder, const Location location, + const ComparisonPredicate predicate, const llvm::APInt& rhs, + const llvm::function_ref loadBit) { + auto one = arith::ConstantIntOp::create(builder, location, 1, 1); + Value equal = one; + Value less; + if (predicate != ComparisonPredicate::Equal && + predicate != ComparisonPredicate::NotEqual) { + less = arith::ConstantIntOp::create(builder, location, 0, 1); + } + for (int64_t index = static_cast(rhs.getBitWidth()) - 1; index >= 0; + --index) { + auto bit = loadBit(index); + Value matches = bit; + if (!rhs[static_cast(index)]) { + matches = arith::XOrIOp::create(builder, location, bit, one); + } else if (less) { + auto lower = arith::XOrIOp::create(builder, location, bit, one); + auto firstDifference = + arith::AndIOp::create(builder, location, equal, lower); + less = arith::OrIOp::create(builder, location, less, firstDifference); + } + equal = arith::AndIOp::create(builder, location, equal, matches); + } + switch (predicate) { + case ComparisonPredicate::Equal: + return equal; + case ComparisonPredicate::NotEqual: + return arith::XOrIOp::create(builder, location, equal, one); + case ComparisonPredicate::Less: + return less; + case ComparisonPredicate::LessEqual: + return arith::OrIOp::create(builder, location, less, equal); + case ComparisonPredicate::Greater: { + auto lessOrEqual = arith::OrIOp::create(builder, location, less, equal); + return arith::XOrIOp::create(builder, location, lessOrEqual, one); + } + case ComparisonPredicate::GreaterEqual: + return arith::XOrIOp::create(builder, location, less, one); + } + llvm_unreachable("unknown CBit comparison predicate"); +} + void LoadOp::getCanonicalizationPatterns(RewritePatternSet& results, MLIRContext* context) { results.add(context); } +void CompareOp::getCanonicalizationPatterns(RewritePatternSet& results, + MLIRContext* context) { + results.add(context); +} + LogicalResult StoreOp::verify() { return verifyIndex(getOperation(), getReg(), getIndex()); } diff --git a/mlir/lib/Dialect/QC/IR/Modifiers/ModifierUtils.cpp b/mlir/lib/Dialect/QC/IR/Modifiers/ModifierUtils.cpp index fd8d89c8e7..7d96a687cb 100644 --- a/mlir/lib/Dialect/QC/IR/Modifiers/ModifierUtils.cpp +++ b/mlir/lib/Dialect/QC/IR/Modifiers/ModifierUtils.cpp @@ -34,9 +34,9 @@ namespace mlir::qc::detail { LogicalResult verifyModifierBody(Operation* modifierOp, Block& body) { const auto hasNonUnitaryOperation = body.walk([](Operation* operation) { - return isa(operation) + return isa(operation) ? WalkResult::interrupt() : WalkResult::advance(); }) diff --git a/mlir/lib/Dialect/QC/Translation/OpenQASMToQCEmitter.cpp b/mlir/lib/Dialect/QC/Translation/OpenQASMToQCEmitter.cpp index c0c2de6659..966f8c59e6 100644 --- a/mlir/lib/Dialect/QC/Translation/OpenQASMToQCEmitter.cpp +++ b/mlir/lib/Dialect/QC/Translation/OpenQASMToQCEmitter.cpp @@ -572,6 +572,9 @@ class OpenQASMToQCEmitter { return chargeDynamicBitRead(condition.bit, multiplicity, projectedEmission, source); } + if (condition.kind == frontend::ConditionKind::RegisterComparison) { + return chargeScaledEmission(1, multiplicity, projectedEmission, source); + } if (condition.kind == frontend::ConditionKind::Comparison) { return chargeExpressionEmission(condition.comparisonLhs, multiplicity, projectedEmission, source) && @@ -1871,6 +1874,25 @@ class OpenQASMToQCEmitter { return builder.loadClassicalBit(reg, registerIndex.getResult()); } + [[nodiscard]] static cbit::ComparisonPredicate + registerPredicate(const frontend::ComparisonKind comparison) { + switch (comparison) { + case frontend::ComparisonKind::Equal: + return cbit::ComparisonPredicate::Equal; + case frontend::ComparisonKind::NotEqual: + return cbit::ComparisonPredicate::NotEqual; + case frontend::ComparisonKind::Less: + return cbit::ComparisonPredicate::Less; + case frontend::ComparisonKind::LessEqual: + return cbit::ComparisonPredicate::LessEqual; + case frontend::ComparisonKind::Greater: + return cbit::ComparisonPredicate::Greater; + case frontend::ComparisonKind::GreaterEqual: + return cbit::ComparisonPredicate::GreaterEqual; + } + llvm_unreachable("unknown register comparison"); + } + [[nodiscard]] Value emitComparison(const frontend::ConditionExpression& condition, ValueRange gateParameters) { @@ -1942,6 +1964,16 @@ class OpenQASMToQCEmitter { return emitQubitOperation( condition.measurement, gateQubits, [&](Value qubit) { return builder.measure(qubit); }); + case frontend::ConditionKind::RegisterComparison: { + auto reg = classicalRegisters.at(condition.reg); + assert(reg && "semantic analysis must declare bit registers before use"); + auto rhs = builder.getIntegerAttr( + builder.getIntegerType(condition.expected.getBitWidth()), + condition.expected); + return cbit::CompareOp::create(builder, builder.getI1Type(), + registerPredicate(condition.comparison), + reg, rhs); + } case frontend::ConditionKind::Not: return arith::XOrIOp::create( builder, emitCondition(condition.lhs, gateParameters, gateQubits), diff --git a/mlir/lib/Dialect/QC/Translation/TranslateQCToOpenQASM3.cpp b/mlir/lib/Dialect/QC/Translation/TranslateQCToOpenQASM3.cpp index 7b6815948b..01808553f0 100644 --- a/mlir/lib/Dialect/QC/Translation/TranslateQCToOpenQASM3.cpp +++ b/mlir/lib/Dialect/QC/Translation/TranslateQCToOpenQASM3.cpp @@ -21,6 +21,7 @@ #include #include +#include #include #include #include @@ -224,6 +225,79 @@ class OpenQASMEmitter { return uniqueName("q", nextQubit); } + [[nodiscard]] static bool storesToRegister(Operation& operation, Value reg) { + return operation + .walk([&](cbit::StoreOp store) { + return store.getReg() == reg ? WalkResult::interrupt() + : WalkResult::advance(); + }) + .wasInterrupted(); + } + + [[nodiscard]] static bool + hasInterveningRegisterWrite(cbit::CompareOp comparison, Operation& consumer) { + Operation* anchor = &consumer; + Block* block = consumer.getBlock(); + Block* comparisonBlock = comparison->getBlock(); + while (block != comparisonBlock) { + for (Operation& operation : *block) { + if (&operation == anchor) { + break; + } + if (storesToRegister(operation, comparison.getReg())) { + return true; + } + } + Operation* parent = block->getParentOp(); + if (parent == nullptr) { + return true; + } + if (isa(parent) && + storesToRegister(*parent, comparison.getReg())) { + return true; + } + anchor = parent; + block = parent->getBlock(); + } + + for (Operation* operation = comparison->getNextNode(); operation != anchor; + operation = operation->getNextNode()) { + if (operation == nullptr || + storesToRegister(*operation, comparison.getReg())) { + return true; + } + } + return false; + } + + [[nodiscard]] LogicalResult validateRegisterComparisonSnapshots() { + const auto walkResult = function.walk([&](cbit::CompareOp comparison) { + SmallVector consumers; + llvm::append_range(consumers, comparison.getResult().getUsers()); + DenseSet visited; + while (!consumers.empty()) { + Operation* consumer = consumers.pop_back_val(); + if (!visited.insert(consumer).second) { + continue; + } + if (hasInterveningRegisterWrite(comparison, *consumer)) { + std::ignore = + fail(comparison, + "register comparison crosses an intervening register write"); + return WalkResult::interrupt(); + } + if (!isInlineExpressionOperation(*consumer)) { + continue; + } + for (Value result : consumer->getResults()) { + llvm::append_range(consumers, result.getUsers()); + } + } + return WalkResult::advance(); + }); + return walkResult.wasInterrupted() ? failure() : success(); + } + [[nodiscard]] LogicalResult preflight() { SmallVector functions(moduleOp.getOps()); if (functions.size() != 1) { @@ -261,7 +335,7 @@ class OpenQASMEmitter { "scope"); } } - return success(); + return validateRegisterComparisonSnapshots(); } [[nodiscard]] LogicalResult collectProgramShape() { @@ -303,8 +377,8 @@ class OpenQASMEmitter { if (auto alloc = dyn_cast(&operation)) { const auto type = alloc.getResult().getType(); const auto width = type.getWidth(); - if (width <= 0 || static_cast(width) > - MAX_CLASSICAL_BITS - numClassicalBits) { + if (width <= 0 || + std::cmp_greater(width, MAX_CLASSICAL_BITS - numClassicalBits)) { return fail(alloc, "total classical register width exceeds the " "supported limit of " + Twine(MAX_CLASSICAL_BITS) + " bits"); @@ -498,7 +572,8 @@ class OpenQASMEmitter { [[nodiscard]] static bool isInlineExpressionOperation(Operation& operation) { const auto name = operation.getName().getStringRef(); - return isa(&operation) || + return isa(&operation) || !binaryOperator(name).empty() || name == "arith.negf" || name == "arith.remf" || isScalarCast(name) || !mathFunction(name).empty(); @@ -588,6 +663,37 @@ class OpenQASMEmitter { if (auto load = value.getDefiningOp()) { return emitBitReference(load.getReg(), load.getIndex()); } + if (auto comparison = value.getDefiningOp()) { + const auto resource = resources.find(comparison.getReg()); + if (resource == resources.end() || + resource->second.kind != ResourceKind::Bit) { + return failExpression(value, + "register comparison refers to unsupported " + "storage"); + } + const auto* predicate = [&] { + switch (comparison.getPredicate()) { + case cbit::ComparisonPredicate::Equal: + return "=="; + case cbit::ComparisonPredicate::NotEqual: + return "!="; + case cbit::ComparisonPredicate::Less: + return "<"; + case cbit::ComparisonPredicate::LessEqual: + return "<="; + case cbit::ComparisonPredicate::Greater: + return ">"; + case cbit::ComparisonPredicate::GreaterEqual: + return ">="; + } + llvm_unreachable("unknown CBit comparison predicate"); + }(); + llvm::SmallString<32> rhs; + comparison.getRhs().toString(rhs, 10, false); + return (Twine("(") + resource->second.name + " " + predicate + " " + rhs + + ")") + .str(); + } auto* operation = value.getDefiningOp(); if (operation == nullptr) { return failExpression(value, "unmapped block argument"); @@ -936,8 +1042,8 @@ class OpenQASMEmitter { return fail(whileOp, "scf.while loop-carried values are not supported"); } for (Operation& operation : before.without_terminator()) { - if (auto load = dyn_cast(operation)) { - if (failed(emitExpression(load.getResult()))) { + if (isa(operation)) { + if (failed(emitExpression(operation.getResult(0)))) { return failure(); } continue; diff --git a/mlir/lib/Dialect/QCO/IR/Modifiers/ModifierUtils.cpp b/mlir/lib/Dialect/QCO/IR/Modifiers/ModifierUtils.cpp index bcae0c576e..9e0d5db824 100644 --- a/mlir/lib/Dialect/QCO/IR/Modifiers/ModifierUtils.cpp +++ b/mlir/lib/Dialect/QCO/IR/Modifiers/ModifierUtils.cpp @@ -33,8 +33,9 @@ namespace mlir::qco::detail { LogicalResult verifyModifierBody(Operation* modifierOp, Block& body) { const auto hasNonUnitaryOperation = body.walk([](Operation* operation) { - return isa(operation) + return isa(operation) ? WalkResult::interrupt() : WalkResult::advance(); }) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 4a8cfb300f..4655a76090 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -360,7 +360,7 @@ static LogicalResult applyUnitaryMatrix(UnitaryOpInterface unitary, } ArrayRef wires = *wiresOr; if (wires.size() >= 63 || - local.rows() != static_cast(size_t{1} << wires.size())) { + std::cmp_not_equal(local.rows(), uint64_t{1} << wires.size())) { return unitary.emitError() << "unitary matrix dimension does not match its target count"; } @@ -655,6 +655,46 @@ static LogicalResult loadRegister(cbit::LoadOp load, ClassicalEnv& classical) { classical); } +static LogicalResult compareRegister(cbit::CompareOp compare, + ClassicalEnv& classical) { + const auto regIt = classical.registers.find(compare.getReg()); + if (regIt == classical.registers.end()) { + return compare.emitError() + << "CBit register is not mapped for QCO DD simulation"; + } + llvm::APInt actual(compare.getRhs().getBitWidth(), 0); + for (const auto [index, cell] : llvm::enumerate(*regIt->second)) { + if (cell.deferredWire && classical.deferredMeasurementUse != nullptr) { + *classical.deferredMeasurementUse = compare.getOperation(); + return failure(); + } + if (!cell.value) { + return compare.emitError() + << "read from an undefined CBit register element"; + } + actual.setBitVal(static_cast(index), *cell.value); + } + const auto result = [&] { + switch (compare.getPredicate()) { + case cbit::ComparisonPredicate::Equal: + return actual.eq(compare.getRhs()); + case cbit::ComparisonPredicate::NotEqual: + return actual.ne(compare.getRhs()); + case cbit::ComparisonPredicate::Less: + return actual.ult(compare.getRhs()); + case cbit::ComparisonPredicate::LessEqual: + return actual.ule(compare.getRhs()); + case cbit::ComparisonPredicate::Greater: + return actual.ugt(compare.getRhs()); + case cbit::ComparisonPredicate::GreaterEqual: + return actual.uge(compare.getRhs()); + } + llvm_unreachable("unknown CBit comparison predicate"); + }(); + return bindInteger(compare.getResult(), + llvm::APInt(1, static_cast(result)), classical); +} + static FailureOr lookupMemRefSlot(Value memref, ValueRange indices, ClassicalEnv& classical, Operation* op) { @@ -1282,6 +1322,9 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { .template Case([&](cbit::LoadOp load) { return loadRegister(load, *walk.classical); }) + .template Case([&](cbit::CompareOp compare) { + return compareRegister(compare, *walk.classical); + }) .template Case([&](cbit::StoreOp store) { return storeRegister(store, *walk.classical); }) diff --git a/mlir/lib/Target/OpenQASM/OpenQASMSemantics.cpp b/mlir/lib/Target/OpenQASM/OpenQASMSemantics.cpp index 8de5d8f5ea..adb3564dd1 100644 --- a/mlir/lib/Target/OpenQASM/OpenQASMSemantics.cpp +++ b/mlir/lib/Target/OpenQASM/OpenQASMSemantics.cpp @@ -3711,7 +3711,7 @@ class SemanticAnalyzer { const auto* lhsSymbol = lhsSyntax.kind == Expr::Kind::Identifier ? lookup(lhsSyntax.identifier) : nullptr; - if (program.openQASM2 && condition.kind == Expr::Kind::Equal && + if ((!program.openQASM2 || condition.kind == Expr::Kind::Equal) && lhsSymbol != nullptr && lhsSymbol->kind == SymbolKind::Register && program.registers[lhsSymbol->id].kind == RegisterKind::Bit && isConstantExpression(*condition.rhs)) { @@ -3719,8 +3719,13 @@ class SemanticAnalyzer { MQT_OQ3_TRY_ASSIGN(bits, resolveBits({.location = lhsSyntax.location, .identifier = lhsSyntax.identifier})); - // OpenQASM 2 classical bits default to 0, so partially written - // registers are valid in `if (c == k)` (e.g. mid-circuit feedback). + if (!program.openQASM2) { + for (const auto& bit : bits) { + if (failed(ensureBitInitialized(bit, condition.location))) { + return failure(); + } + } + } llvm::APInt expectedBits; if (rhsSyntax.kind == Expr::Kind::Int && !rhsSyntax.wideInteger.empty()) { @@ -3740,7 +3745,7 @@ class SemanticAnalyzer { std::get(expected.value) < 0)) { return fail( condition.location, - "OpenQASM 2 register conditions require an unsigned integer"); + "classical register conditions require an unsigned integer"); } const auto expectedValue = expected.type == ScalarType::Uint @@ -3749,39 +3754,41 @@ class SemanticAnalyzer { expectedBits = llvm::APInt(/*numBits=*/64, expectedValue); } if (expectedBits.getActiveBits() > bits.size()) { - // Value cannot equal the register contents. + const bool result = condition.kind == Expr::Kind::NotEqual || + condition.kind == Expr::Kind::Less || + condition.kind == Expr::Kind::LessEqual; return addCondition( {.kind = ConditionKind::Literal, .location = getSourceLocation(condition.location), - .literal = false}); + .literal = result}); } if (expectedBits.getBitWidth() < bits.size()) { expectedBits = expectedBits.zext(static_cast(bits.size())); } else if (expectedBits.getBitWidth() > bits.size()) { expectedBits = expectedBits.trunc(static_cast(bits.size())); } - auto result = - addCondition({.kind = ConditionKind::Literal, - .location = getSourceLocation(condition.location), - .literal = true}); - for (const auto [index, bit] : llvm::enumerate(bits)) { - auto bitCondition = - addCondition({.kind = ConditionKind::Bit, - .location = getSourceLocation(condition.location), - .bit = bit}); - if (!expectedBits[index]) { - bitCondition = - addCondition({.kind = ConditionKind::Not, - .location = getSourceLocation(condition.location), - .lhs = bitCondition}); - } - result = - addCondition({.kind = ConditionKind::And, - .location = getSourceLocation(condition.location), - .lhs = result, - .rhs = bitCondition}); - } - return result; + return addCondition({.kind = ConditionKind::RegisterComparison, + .location = getSourceLocation(condition.location), + .reg = lhsSymbol->id, + .expected = std::move(expectedBits), + .comparison = [&] { + switch (condition.kind) { + case Expr::Kind::Equal: + return ComparisonKind::Equal; + case Expr::Kind::NotEqual: + return ComparisonKind::NotEqual; + case Expr::Kind::Less: + return ComparisonKind::Less; + case Expr::Kind::LessEqual: + return ComparisonKind::LessEqual; + case Expr::Kind::Greater: + return ComparisonKind::Greater; + case Expr::Kind::GreaterEqual: + return ComparisonKind::GreaterEqual; + default: + llvm_unreachable("not a comparison"); + } + }()}); } typed.kind = ConditionKind::Comparison; MQT_OQ3_TRY_ASSIGN(comparisonLhs, analyzeExpression(*condition.lhs)); diff --git a/mlir/unittests/Conversion/CBitToMemRef/test_cbit_to_memref.cpp b/mlir/unittests/Conversion/CBitToMemRef/test_cbit_to_memref.cpp index a48ff2acd2..8ed0ac25e2 100644 --- a/mlir/unittests/Conversion/CBitToMemRef/test_cbit_to_memref.cpp +++ b/mlir/unittests/Conversion/CBitToMemRef/test_cbit_to_memref.cpp @@ -134,6 +134,34 @@ TEST_F(CBitToMemRefTest, LargeZeroInitializationProducesBoundedIR) { EXPECT_EQ(stores, 1); } +TEST_F(CBitToMemRefTest, LowersRegisterComparisons) { + auto moduleOp = convert(R"mlir( + module { + func.func @main() -> (i1, i1, i1, i1, i1, i1) { + %reg = cbit.alloc(#cbit.init) : !cbit.reg<3> + %eq = cbit.cmp eq, %reg, 5 : i3 : !cbit.reg<3> + %ne = cbit.cmp ne, %reg, 5 : i3 : !cbit.reg<3> + %ult = cbit.cmp ult, %reg, 5 : i3 : !cbit.reg<3> + %ule = cbit.cmp ule, %reg, 5 : i3 : !cbit.reg<3> + %ugt = cbit.cmp ugt, %reg, 5 : i3 : !cbit.reg<3> + %uge = cbit.cmp uge, %reg, 5 : i3 : !cbit.reg<3> + return %eq, %ne, %ult, %ule, %ugt, %uge : i1, i1, i1, i1, i1, i1 + } + } + )mlir"); + ASSERT_TRUE(moduleOp); + EXPECT_TRUE(succeeded(verify(*moduleOp))); + + bool containsCBit = false; + moduleOp->walk([&](Operation* op) { + containsCBit |= op->getDialect() == context->getLoadedDialect("cbit"); + }); + EXPECT_FALSE(containsCBit); + size_t loads = 0; + moduleOp->walk([&](memref::LoadOp) { ++loads; }); + EXPECT_EQ(loads, 18); +} + TEST_F(CBitToMemRefTest, ConvertsFunctionSignaturesCallsAndReturns) { auto moduleOp = convert(R"mlir( module { diff --git a/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp b/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp index 05494a9418..31f4783e52 100644 --- a/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp +++ b/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp @@ -1200,13 +1200,15 @@ TEST_F(QCToQCORegressionTest, } namespace { -enum class CBitModifierBodyOp : std::uint8_t { Alloc, Load, Store }; +enum class CBitModifierBodyOp : std::uint8_t { Alloc, Compare, Load, Store }; } // namespace static StringRef cbitOperationName(CBitModifierBodyOp operation) { switch (operation) { case CBitModifierBodyOp::Alloc: return "cbit.alloc"; + case CBitModifierBodyOp::Compare: + return "cbit.cmp"; case CBitModifierBodyOp::Load: return "cbit.load"; case CBitModifierBodyOp::Store: @@ -1233,6 +1235,11 @@ buildInvalidCBitModifierProgram(MLIRContext* context, cbit::RegisterType::get(builder.getContext(), 1), cbit::Initialization::Zero); break; + case CBitModifierBodyOp::Compare: + cbit::CompareOp::create(builder, builder.getI1Type(), + cbit::ComparisonPredicate::Equal, reg, + builder.getIntegerAttr(builder.getI1Type(), 0)); + break; case CBitModifierBodyOp::Load: cbit::LoadOp::create(builder, builder.getI1Type(), reg, index.getResult()); @@ -1262,9 +1269,9 @@ TEST_F(QCToQCORegressionTest, PreflightRejectsEveryCBitOperationInEveryModifier) { constexpr std::array modifiers{ModifierKind::Inv, ModifierKind::Ctrl, ModifierKind::Pow}; - constexpr std::array operations{CBitModifierBodyOp::Alloc, - CBitModifierBodyOp::Load, - CBitModifierBodyOp::Store}; + constexpr std::array operations{ + CBitModifierBodyOp::Alloc, CBitModifierBodyOp::Compare, + CBitModifierBodyOp::Load, CBitModifierBodyOp::Store}; for (const auto modifier : modifiers) { for (const auto operation : operations) { diff --git a/mlir/unittests/Conversion/QCToQIR/QCToQIRAdaptive/test_qc_to_qir_adaptive.cpp b/mlir/unittests/Conversion/QCToQIR/QCToQIRAdaptive/test_qc_to_qir_adaptive.cpp index 8dc77847ac..2947de4831 100644 --- a/mlir/unittests/Conversion/QCToQIR/QCToQIRAdaptive/test_qc_to_qir_adaptive.cpp +++ b/mlir/unittests/Conversion/QCToQIR/QCToQIRAdaptive/test_qc_to_qir_adaptive.cpp @@ -11,6 +11,7 @@ #include "Support/IRVerification.h" #include "TestCaseUtils.h" #include "mlir/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.h" +#include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/MQT/Transforms/Passes.h" #include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" #include "mlir/Dialect/QC/IR/QCDialect.h" @@ -35,6 +36,7 @@ #include #include #include +#include #include #include #include @@ -256,6 +258,109 @@ TEST(QCToQIRAdaptiveNativeTest, LowersZeroInitializedClassicalControlRegister) { EXPECT_FALSE(module->lookupSymbol(qir::QIR_READ_RESULT)); } +TEST(QCToQIRAdaptiveNativeTest, LowersClassicalRegisterComparison) { + MLIRContext context; + context + .loadDialect(); + qc::QCProgramBuilder builder(&context); + builder.initialize(); + auto q = builder.allocQubit(); + auto c = builder.allocClassicalBitRegister(3); + builder.measure(q, c, 0); + auto rhs = builder.getIntegerAttr(builder.getIntegerType(3), 2); + auto comparison = cbit::CompareOp::create( + builder, builder.getI1Type(), cbit::ComparisonPredicate::Less, c, rhs); + builder.scfIf(comparison, [&] { builder.x(q); }); + auto module = builder.finalize(); + ASSERT_TRUE(module); + ASSERT_TRUE(succeeded(verify(*module))); + ASSERT_TRUE(succeeded(runQCToQIRAdaptiveConversion(*module))); + EXPECT_TRUE(succeeded(verify(*module))); + bool retainsComparison = false; + module->walk([&](cbit::CompareOp) { retainsComparison = true; }); + EXPECT_FALSE(retainsComparison); +} + +TEST(QCToQIRAdaptiveNativeTest, RejectsMixedClassicalRegisterRepresentations) { + MLIRContext context; + context.loadDialect(); + auto module = parseSourceString(R"mlir( + module { + func.func @main() -> (i1, i1, !cbit.reg<1>) attributes {mqt.entry_point} { + %true = arith.constant true + %c0 = arith.constant 0 : index + %returned = cbit.alloc(#cbit.init) : !cbit.reg<1> + %local = cbit.alloc(#cbit.init) : !cbit.reg<1> + %selected = scf.if %true -> (!cbit.reg<1>) { + scf.yield %returned : !cbit.reg<1> + } else { + scf.yield %local : !cbit.reg<1> + } + %bit = cbit.load %selected[%c0] : !cbit.reg<1> + %matches = cbit.cmp eq, %selected, 0 : i1 : !cbit.reg<1> + return %bit, %matches, %returned : i1, i1, !cbit.reg<1> + } + } + )mlir", + &context); + ASSERT_TRUE(module); + ASSERT_TRUE(succeeded(verify(*module))); + + size_t mixedRepresentationDiagnostics = 0; + const ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + std::string message; + llvm::raw_string_ostream stream(message); + diagnostic.print(stream); + mixedRepresentationDiagnostics += StringRef(message).contains( + "adaptive QIR conversion cannot merge returned and local CBit " + "registers"); + return success(); + }); + EXPECT_TRUE(failed(runQCToQIRAdaptiveConversionSimple(*module))); + EXPECT_EQ(mixedRepresentationDiagnostics, 2); +} + +TEST(QCToQIRAdaptiveNativeTest, LowersReturnedRegisterMerge) { + MLIRContext context; + context.loadDialect(); + auto module = parseSourceString(R"mlir( + module { + func.func @main() -> (i1, i1, !cbit.reg<1>, !cbit.reg<1>) + attributes {mqt.entry_point} { + %true = arith.constant true + %c0 = arith.constant 0 : index + %first = cbit.alloc(#cbit.init) : !cbit.reg<1> + %second = cbit.alloc(#cbit.init) : !cbit.reg<1> + %selected = scf.if %true -> (!cbit.reg<1>) { + scf.yield %first : !cbit.reg<1> + } else { + scf.yield %second : !cbit.reg<1> + } + %bit = cbit.load %selected[%c0] : !cbit.reg<1> + %matches = cbit.cmp eq, %selected, 0 : i1 : !cbit.reg<1> + return %bit, %matches, %first, %second + : i1, i1, !cbit.reg<1>, !cbit.reg<1> + } + } + )mlir", + &context); + ASSERT_TRUE(module); + ASSERT_TRUE(succeeded(verify(*module))); + EXPECT_TRUE(succeeded(runQCToQIRAdaptiveConversionSimple(*module))); + EXPECT_TRUE(succeeded(verify(*module))); + size_t resultReads = 0; + module->walk([&](LLVM::CallOp call) { + resultReads += call.getCallee() == qir::QIR_READ_RESULT; + }); + EXPECT_EQ(resultReads, 2); +} + TEST(QCToQIRAdaptiveNativeTest, RejectsMultipleRegisterDestinations) { MLIRContext context; context diff --git a/mlir/unittests/Conversion/QCToQIR/QCToQIRBase/test_qc_to_qir_base.cpp b/mlir/unittests/Conversion/QCToQIR/QCToQIRBase/test_qc_to_qir_base.cpp index 41bfced7c1..e92df59f35 100644 --- a/mlir/unittests/Conversion/QCToQIR/QCToQIRBase/test_qc_to_qir_base.cpp +++ b/mlir/unittests/Conversion/QCToQIR/QCToQIRBase/test_qc_to_qir_base.cpp @@ -11,6 +11,8 @@ #include "Support/IRVerification.h" #include "TestCaseUtils.h" #include "mlir/Conversion/QCToQIR/QIRBase/QCToQIRBase.h" +#include "mlir/Dialect/CBit/IR/CBitDialect.h" +#include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/MQT/Transforms/Passes.h" #include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" #include "mlir/Dialect/QC/IR/QCDialect.h" @@ -157,6 +159,32 @@ TEST(QCToQIRBaseNativeTest, RejectsMultiBlockEntryFunctionWithoutMutation) { EXPECT_EQ(entryPoint.getBlocks().size(), 2); } +TEST(QCToQIRBaseNativeTest, RejectsClassicalRegisterComparisons) { + MLIRContext context; + context.loadDialect(); + qc::QCProgramBuilder builder(&context); + builder.initialize(); + auto reg = builder.allocClassicalBitRegister(1); + auto rhs = builder.getIntegerAttr(builder.getIntegerType(1), 0); + (void)cbit::CompareOp::create(builder, builder.getI1Type(), + cbit::ComparisonPredicate::Equal, reg, rhs); + auto module = builder.finalize(); + ASSERT_TRUE(module); + + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + std::string message; + llvm::raw_string_ostream stream(message); + diagnostic.print(stream); + sawExpectedDiagnostic |= StringRef(message).contains( + "QIR Base Profile does not support classical-register comparisons"); + return success(); + }); + EXPECT_TRUE(failed(runQCToQIRBaseConversion(*module))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + TEST(QCToQIRBaseNativeTest, ControlledBarrierDoesNotControlFollowingGate) { expectFollowingXIsUncontrolled( [](qc::QCProgramBuilder& builder, Value control, Value target) { diff --git a/mlir/unittests/Dialect/CBit/IR/test_cbit_ir.cpp b/mlir/unittests/Dialect/CBit/IR/test_cbit_ir.cpp index 6bdb47a4f5..6baa759b85 100644 --- a/mlir/unittests/Dialect/CBit/IR/test_cbit_ir.cpp +++ b/mlir/unittests/Dialect/CBit/IR/test_cbit_ir.cpp @@ -66,6 +66,7 @@ TEST_F(CBitIRTest, ParsesAndPrintsRegisterOperations) { %reg = cbit.alloc(#cbit.init) {mqt.register_name = "c"} : !cbit.reg<2> cbit.store %false, %reg[%c0] : !cbit.reg<2> %bit = cbit.load %reg[%c0] : !cbit.reg<2> + %matches = cbit.cmp eq, %reg, 1 : i2 : !cbit.reg<2> return %reg : !cbit.reg<2> } } @@ -84,6 +85,31 @@ TEST_F(CBitIRTest, ParsesAndPrintsRegisterOperations) { EXPECT_NE(printed.find("!cbit.reg<2>"), std::string::npos); EXPECT_NE(printed.find("cbit.store"), std::string::npos); EXPECT_NE(printed.find("cbit.load"), std::string::npos); + EXPECT_NE(printed.find("cbit.cmp eq"), std::string::npos); +} + +TEST_F(CBitIRTest, RejectsComparisonWidthMismatch) { + EXPECT_FALSE(parse(R"mlir( + module { + func.func @main() { + %reg = cbit.alloc(#cbit.init) : !cbit.reg<2> + %matches = cbit.cmp eq, %reg, 1 : i3 : !cbit.reg<2> + return + } + } + )mlir")); +} + +TEST_F(CBitIRTest, RejectsUnsupportedComparisonWidth) { + EXPECT_FALSE(parse(R"mlir( + module { + func.func @main() { + %reg = cbit.alloc(#cbit.init) : !cbit.reg<4294967297> + %matches = cbit.cmp eq, %reg, 0 : i1 : !cbit.reg<4294967297> + return + } + } + )mlir")); } TEST_F(CBitIRTest, RejectsNonPositiveRegisterWidth) { @@ -146,6 +172,7 @@ TEST_F(CBitIRTest, ReportsMemoryEffects) { %reg = cbit.alloc(#cbit.init) : !cbit.reg<1> cbit.store %false, %reg[%c0] : !cbit.reg<1> %bit = cbit.load %reg[%c0] : !cbit.reg<1> + %matches = cbit.cmp eq, %reg, 0 : i1 : !cbit.reg<1> return } } @@ -153,13 +180,16 @@ TEST_F(CBitIRTest, ReportsMemoryEffects) { ASSERT_TRUE(moduleOp); cbit::AllocOp alloc; + cbit::CompareOp compare; cbit::LoadOp load; cbit::StoreOp store; moduleOp->walk([&](cbit::AllocOp op) { alloc = op; }); + moduleOp->walk([&](cbit::CompareOp op) { compare = op; }); moduleOp->walk([&](cbit::LoadOp op) { load = op; }); moduleOp->walk([&](cbit::StoreOp op) { store = op; }); ASSERT_NE(alloc.getOperation(), nullptr); + ASSERT_NE(compare.getOperation(), nullptr); ASSERT_NE(load.getOperation(), nullptr); ASSERT_NE(store.getOperation(), nullptr); @@ -174,6 +204,12 @@ TEST_F(CBitIRTest, ReportsMemoryEffects) { EXPECT_TRUE(isa(effects.front().getEffect())); EXPECT_EQ(effects.front().getValue(), load.getReg()); + effects.clear(); + compare.getEffects(effects); + ASSERT_EQ(effects.size(), 1); + EXPECT_TRUE(isa(effects.front().getEffect())); + EXPECT_EQ(effects.front().getValue(), compare.getReg()); + effects.clear(); store.getEffects(effects); ASSERT_EQ(effects.size(), 1); @@ -184,15 +220,16 @@ TEST_F(CBitIRTest, ReportsMemoryEffects) { TEST_F(CBitIRTest, ForwardsStraightLineStoresAndZeroInitialization) { auto moduleOp = parse(R"mlir( module { - func.func @main() -> (i1, i1) { + func.func @main() -> (i1, i1, i1) { %c0 = arith.constant 0 : index %c1 = arith.constant 1 : index %true = arith.constant true %reg = cbit.alloc(#cbit.init) : !cbit.reg<2> %zero = cbit.load %reg[%c0] : !cbit.reg<2> + %matches = cbit.cmp eq, %reg, 0 : i2 : !cbit.reg<2> cbit.store %true, %reg[%c1] : !cbit.reg<2> %stored = cbit.load %reg[%c1] : !cbit.reg<2> - return %zero, %stored : i1, i1 + return %zero, %stored, %matches : i1, i1, i1 } } )mlir"); @@ -209,24 +246,29 @@ TEST_F(CBitIRTest, ForwardsStraightLineStoresAndZeroInitialization) { moduleOp->print(canonicalizedStream); APInt zero; APInt stored; + APInt matches; EXPECT_TRUE(matchPattern(returnOp.getOperand(0), m_ConstantInt(&zero))) << canonicalized; EXPECT_TRUE(matchPattern(returnOp.getOperand(1), m_ConstantInt(&stored))) << canonicalized; + EXPECT_TRUE(matchPattern(returnOp.getOperand(2), m_ConstantInt(&matches))) + << canonicalized; EXPECT_TRUE(zero.isZero()); EXPECT_TRUE(stored.isOne()); + EXPECT_TRUE(matches.isOne()); } TEST_F(CBitIRTest, DoesNotForwardAcrossAnAmbiguousStore) { auto moduleOp = parse(R"mlir( module { - func.func @main(%dynamic: index) -> i1 { + func.func @main(%dynamic: index) -> (i1, i1) { %c0 = arith.constant 0 : index %true = arith.constant true %reg = cbit.alloc(#cbit.init) : !cbit.reg<2> cbit.store %true, %reg[%dynamic] : !cbit.reg<2> %value = cbit.load %reg[%c0] : !cbit.reg<2> - return %value : i1 + %matches = cbit.cmp eq, %reg, 0 : i2 : !cbit.reg<2> + return %value, %matches : i1, i1 } } )mlir"); @@ -239,5 +281,6 @@ TEST_F(CBitIRTest, DoesNotForwardAcrossAnAmbiguousStore) { auto funcOp = *moduleOp->getOps().begin(); auto returnOp = *funcOp.getOps().begin(); EXPECT_TRUE(returnOp.getOperand(0).getDefiningOp()); + EXPECT_TRUE(returnOp.getOperand(1).getDefiningOp()); } } // namespace diff --git a/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp b/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp index fa17cf9476..f3561243bd 100644 --- a/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp +++ b/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp @@ -622,6 +622,7 @@ enum class ForbiddenModifierBodyOp : std::uint8_t { QubitRegisterLoad, QubitRegisterStore, CBitAlloc, + CBitCompare, CBitLoad, CBitStore }; @@ -656,6 +657,8 @@ static StringRef forbiddenOperationName(ForbiddenModifierBodyOp kind) { return "qubit-register-store"; case ForbiddenModifierBodyOp::CBitAlloc: return "cbit.alloc"; + case ForbiddenModifierBodyOp::CBitCompare: + return "cbit.cmp"; case ForbiddenModifierBodyOp::CBitLoad: return "cbit.load"; case ForbiddenModifierBodyOp::CBitStore: @@ -693,6 +696,11 @@ static void emitForbiddenModifierBodyOperation(QCProgramBuilder& builder, cbit::RegisterType::get(builder.getContext(), 1), cbit::Initialization::Zero); return; + case ForbiddenModifierBodyOp::CBitCompare: + cbit::CompareOp::create(builder, builder.getI1Type(), + cbit::ComparisonPredicate::Equal, cbitReg, + builder.getIntegerAttr(builder.getI1Type(), 0)); + return; case ForbiddenModifierBodyOp::CBitLoad: cbit::LoadOp::create(builder, builder.getI1Type(), cbitReg, index); return; @@ -749,6 +757,7 @@ TEST_F(QCTest, ModifiersRecursivelyRejectEveryForbiddenOperation) { ForbiddenModifierBodyOp::QubitRegisterLoad, ForbiddenModifierBodyOp::QubitRegisterStore, ForbiddenModifierBodyOp::CBitAlloc, + ForbiddenModifierBodyOp::CBitCompare, ForbiddenModifierBodyOp::CBitLoad, ForbiddenModifierBodyOp::CBitStore}; diff --git a/mlir/unittests/Dialect/QC/Translation/test_openqasm3_emission.cpp b/mlir/unittests/Dialect/QC/Translation/test_openqasm3_emission.cpp index 5773799a5b..46eac456bd 100644 --- a/mlir/unittests/Dialect/QC/Translation/test_openqasm3_emission.cpp +++ b/mlir/unittests/Dialect/QC/Translation/test_openqasm3_emission.cpp @@ -268,6 +268,109 @@ switch (selector) { << *emitted; } +TEST(OpenQASM3EmissionTest, RoundTripsRegisterComparisonInWhileCondition) { + constexpr llvm::StringLiteral source = R"qasm(OPENQASM 3.1; +include "stdgates.inc"; +qubit q; +bit[1] c = measure q; +while (c == 1) { + c[0] = measure q; +} +)qasm"; + MLIRContext context; + auto moduleOp = qc::translateQASM3ToQC(source, &context); + ASSERT_TRUE(moduleOp); + + auto emitted = qc::translateQCToOpenQASM3(*moduleOp); + + ASSERT_TRUE(succeeded(emitted)); + EXPECT_NE(emitted->find("while ("), std::string::npos) << *emitted; + EXPECT_NE(emitted->find("c == 1"), std::string::npos) << *emitted; + EXPECT_TRUE(oq3::frontend::analyzeOpenQASM( + *emitted, {.gatePolicy = oq3::frontend::GatePolicy::Strict})) + << *emitted; +} + +TEST(OpenQASM3EmissionTest, EmitsRegisterComparisonsDirectly) { + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func @main() -> !cbit.reg<3> attributes {mqt.entry_point} { + %q = qc.alloc : !qc.qubit + %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} + : !cbit.reg<3> + %eq = cbit.cmp eq, %c, 5 : i3 : !cbit.reg<3> + %ne = cbit.cmp ne, %c, 5 : i3 : !cbit.reg<3> + %ult = cbit.cmp ult, %c, 5 : i3 : !cbit.reg<3> + %ule = cbit.cmp ule, %c, 5 : i3 : !cbit.reg<3> + %ugt = cbit.cmp ugt, %c, 5 : i3 : !cbit.reg<3> + %uge = cbit.cmp uge, %c, 5 : i3 : !cbit.reg<3> + scf.if %eq { + qc.x %q : !qc.qubit + } + scf.if %ne { + qc.x %q : !qc.qubit + } + scf.if %ult { + qc.x %q : !qc.qubit + } + scf.if %ule { + qc.x %q : !qc.qubit + } + scf.if %ugt { + qc.x %q : !qc.qubit + } + scf.if %uge { + qc.x %q : !qc.qubit + } + return %c : !cbit.reg<3> + } +} +)mlir"; + DialectRegistry registry = emissionDialects(); + MLIRContext context(registry); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + + auto emitted = qc::translateQCToOpenQASM3(*moduleOp); + + ASSERT_TRUE(succeeded(emitted)); + for (const auto* comparison : + {"c == 5", "c != 5", "c < 5", "c <= 5", "c > 5", "c >= 5"}) { + EXPECT_NE(emitted->find(comparison), std::string::npos) << *emitted; + } + EXPECT_TRUE(oq3::frontend::analyzeOpenQASM( + *emitted, {.gatePolicy = oq3::frontend::GatePolicy::Strict})) + << *emitted; +} + +TEST(OpenQASM3EmissionTest, RejectsComparisonAfterInterveningRegisterWrite) { + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func @main() -> !cbit.reg<3> attributes {mqt.entry_point} { + %zero = arith.constant 0 : index + %true = arith.constant true + %q = qc.alloc : !qc.qubit + %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} + : !cbit.reg<3> + %condition = cbit.cmp eq, %c, 0 : i3 : !cbit.reg<3> + %false = arith.constant false + %forwarded = arith.xori %condition, %false : i1 + cbit.store %true, %c[%zero] : !cbit.reg<3> + scf.if %forwarded { + qc.x %q : !qc.qubit + } + return %c : !cbit.reg<3> + } +} +)mlir"; + DialectRegistry registry = emissionDialects(); + MLIRContext context(registry); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + + EXPECT_TRUE(failed(qc::translateQCToOpenQASM3(*moduleOp))); +} + TEST(OpenQASM3EmissionTest, EmitsNativeIndexSwitch) { constexpr llvm::StringLiteral source = R"mlir( module { diff --git a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp index f174874843..37c2cde61c 100644 --- a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp +++ b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp @@ -394,6 +394,7 @@ enum class VerifierModifierKind : uint8_t { Inv, Ctrl, Pow }; enum class ForbiddenModifierBodyOp : uint8_t { Measure, CBitAlloc, + CBitCompare, CBitLoad, CBitStore }; @@ -418,6 +419,8 @@ static StringRef forbiddenOperationName(ForbiddenModifierBodyOp kind) { return "measure"; case ForbiddenModifierBodyOp::CBitAlloc: return "cbit.alloc"; + case ForbiddenModifierBodyOp::CBitCompare: + return "cbit.cmp"; case ForbiddenModifierBodyOp::CBitLoad: return "cbit.load"; case ForbiddenModifierBodyOp::CBitStore: @@ -480,6 +483,11 @@ buildInvalidNestedModifierBody(QCOProgramBuilder& builder, builder, cbit::RegisterType::get(builder.getContext(), 1), cbit::Initialization::Zero); break; + case ForbiddenModifierBodyOp::CBitCompare: + cbit::CompareOp::create( + builder, builder.getI1Type(), cbit::ComparisonPredicate::Equal, + cbitReg, builder.getIntegerAttr(builder.getI1Type(), 0)); + break; case ForbiddenModifierBodyOp::CBitLoad: cbit::LoadOp::create(builder, builder.getI1Type(), cbitReg, index.getResult()); @@ -512,7 +520,8 @@ TEST_F(QCOTest, ModifiersRecursivelyRejectNonUnitaryOperations) { VerifierModifierKind::Pow}; constexpr std::array forbiddenOperations{ ForbiddenModifierBodyOp::Measure, ForbiddenModifierBodyOp::CBitAlloc, - ForbiddenModifierBodyOp::CBitLoad, ForbiddenModifierBodyOp::CBitStore}; + ForbiddenModifierBodyOp::CBitCompare, ForbiddenModifierBodyOp::CBitLoad, + ForbiddenModifierBodyOp::CBitStore}; for (const auto modifier : modifiers) { for (const auto forbiddenOperation : forbiddenOperations) { diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 9256119ae7..d1c1f7a9e5 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -16,6 +16,7 @@ #include "dd/StateGeneration.hpp" #include "mlir/Dialect/CBit/IR/CBitAttributes.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" +#include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/QCO/Builder/QCOProgramBuilder.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/Utils/DDFunctionality.h" @@ -39,6 +40,7 @@ #include #include +#include #include #include #include @@ -796,6 +798,33 @@ TEST_F(QCODDFunctionalityTest, SimulateCBitConditionAndMeasurementUpdate) { dd->decRef(one); } +TEST_F(QCODDFunctionalityTest, SimulateCBitRegisterComparisons) { + constexpr std::array comparisons{ + std::pair{cbit::ComparisonPredicate::Equal, false}, + std::pair{cbit::ComparisonPredicate::NotEqual, true}, + std::pair{cbit::ComparisonPredicate::Less, true}, + std::pair{cbit::ComparisonPredicate::LessEqual, true}, + std::pair{cbit::ComparisonPredicate::Greater, false}, + std::pair{cbit::ComparisonPredicate::GreaterEqual, false}, + }; + for (const auto [predicate, expected] : comparisons) { + auto mod = buildModule([&](QCOProgramBuilder& b) { + auto reg = b.allocClassicalBitRegister(2, "c"); + auto rhs = b.getIntegerAttr(b.getIntegerType(2), 1); + auto condition = + cbit::CompareOp::create(b, b.getI1Type(), predicate, reg, rhs); + auto q = b.staticQubit(0); + q = b.qcoIf( + condition, q, [&](Value arg) { return b.x(arg); }, + [&](Value arg) { return arg; }); + b.sink(q); + return b.intConstant(0); + }); + ASSERT_TRUE(mod); + expectSimulatesFromZero(mainFunc(*mod), expected); + } +} + TEST_F(QCODDFunctionalityTest, RejectsUndefinedCBitLoad) { auto mod = buildModule([](QCOProgramBuilder& b) { auto reg = @@ -815,6 +844,28 @@ TEST_F(QCODDFunctionalityTest, RejectsUndefinedCBitLoad) { failed(simulate(mainFunc(*mod), dd::makeZeroState(1, *dd), *dd, rng))); } +TEST_F(QCODDFunctionalityTest, RejectsUndefinedCBitRegisterComparison) { + auto mod = buildModule([](QCOProgramBuilder& b) { + auto reg = + b.allocClassicalBitRegister(1, "c", cbit::Initialization::Undefined); + auto rhs = b.getIntegerAttr(b.getIntegerType(1), 0); + auto condition = cbit::CompareOp::create( + b, b.getI1Type(), cbit::ComparisonPredicate::Equal, reg, rhs); + auto q = b.staticQubit(0); + q = b.qcoIf( + condition, q, [&](Value arg) { return arg; }, + [&](Value arg) { return arg; }); + b.sink(q); + return b.intConstant(0); + }); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(1); + std::mt19937_64 rng(1); + EXPECT_TRUE( + failed(simulate(mainFunc(*mod), dd::makeZeroState(1, *dd), *dd, rng))); +} + TEST_F(QCODDFunctionalityTest, SimulateMeasureFeedsIndexSwitch) { // |1> → measure → index_castui → index_switch case 1 applies X → |0>. auto mod = buildModule([](QCOProgramBuilder& b) { diff --git a/mlir/unittests/Target/OpenQASM/test_openqasm_emitter.cpp b/mlir/unittests/Target/OpenQASM/test_openqasm_emitter.cpp index 871c70245d..841bd0ab70 100644 --- a/mlir/unittests/Target/OpenQASM/test_openqasm_emitter.cpp +++ b/mlir/unittests/Target/OpenQASM/test_openqasm_emitter.cpp @@ -1258,10 +1258,59 @@ if (c == 1) x q[0]; ASSERT_TRUE(moduleOp); ASSERT_TRUE(succeeded(verify(*moduleOp))); size_t conditionals = 0; + size_t comparisons = 0; moduleOp->walk([&](scf::IfOp) { ++conditionals; }); - // The register equality and the source-level branch each short-circuit - // through their own structured conditional. - EXPECT_EQ(conditionals, 2); + moduleOp->walk([&](cbit::CompareOp) { ++comparisons; }); + EXPECT_EQ(conditionals, 1); + EXPECT_EQ(comparisons, 1); +} + +TEST(OpenQASMTargetTest, EmitsAllRegisterComparisonPredicates) { + constexpr llvm::StringLiteral source = R"qasm( +OPENQASM 3.1; +include "stdgates.inc"; +bit[3] c; +c[0] = true; +c[1] = false; +c[2] = true; +qubit q; +if (c == 5) { x q; } +if (c != 5) { x q; } +if (c < 5) { x q; } +if (c <= 5) { x q; } +if (c > 5) { x q; } +if (c >= 5) { x q; } +)qasm"; + + MLIRContext context; + auto moduleOp = qc::translateQASM3ToQC(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + std::array predicates{}; + moduleOp->walk([&](cbit::CompareOp comparison) { + predicates.at(static_cast(comparison.getPredicate())) = true; + EXPECT_EQ(comparison.getRhs(), llvm::APInt(3, 5)); + }); + EXPECT_TRUE(llvm::all_of(predicates, [](const bool value) { return value; })); +} + +TEST(OpenQASMTargetTest, PreservesWideRegisterComparisons) { + constexpr llvm::StringLiteral source = R"qasm( +OPENQASM 2.0; +include "qelib1.inc"; +qreg q[1]; +creg c[65]; +if (c == 18446744073709551616) x q[0]; +)qasm"; + + MLIRContext context; + auto moduleOp = qc::translateQASM3ToQC(source, &context); + ASSERT_TRUE(moduleOp); + cbit::CompareOp comparison; + moduleOp->walk([&](cbit::CompareOp op) { comparison = op; }); + ASSERT_TRUE(comparison); + EXPECT_EQ(comparison.getRhs(), llvm::APInt(65, 1).shl(64)); } TEST(OpenQASMTargetTest, ZeroInitializesUnmeasuredOpenQASM2Registers) { diff --git a/mlir/unittests/Target/OpenQASM/test_openqasm_semantics.cpp b/mlir/unittests/Target/OpenQASM/test_openqasm_semantics.cpp index c1011b47cc..f138e1cee8 100644 --- a/mlir/unittests/Target/OpenQASM/test_openqasm_semantics.cpp +++ b/mlir/unittests/Target/OpenQASM/test_openqasm_semantics.cpp @@ -407,6 +407,12 @@ OPENQASM 3.1; qubit q; bit c; if (c) { x q; } +)qasm"; + constexpr llvm::StringLiteral unmeasuredRegisterCondition = R"qasm( +OPENQASM 3.1; +qubit q; +bit[2] c; +if (c >= 1) { x q; } )qasm"; auto uninitializedOutput = oq3::frontend::analyzeOpenQASM(unmeasuredOutput); @@ -423,6 +429,14 @@ if (c) { x q; } EXPECT_NE(uninitializedCondition.diagnostics.front().message.find( "has not been initialized"), std::string::npos); + + auto uninitializedRegister = + oq3::frontend::analyzeOpenQASM(unmeasuredRegisterCondition); + ASSERT_FALSE(uninitializedRegister); + ASSERT_FALSE(uninitializedRegister.diagnostics.empty()); + EXPECT_NE(uninitializedRegister.diagnostics.front().message.find( + "has not been initialized"), + std::string::npos); } TEST(OpenQASMFrontendTest, RejectsUninitializedScalarOutputs) { @@ -1651,7 +1665,8 @@ if(c==1180591620717411303433) x q[0]; auto analyzed = oq3::frontend::analyzeOpenQASM(source); ASSERT_TRUE(analyzed) << analyzed.diagnostics.front().message; EXPECT_TRUE(llvm::any_of(analyzed.program->conditions, [](const auto& c) { - return c.kind == oq3::frontend::ConditionKind::Bit && c.bit.index == 70; + return c.kind == oq3::frontend::ConditionKind::RegisterComparison && + c.expected[70]; })); } @@ -1669,7 +1684,8 @@ if(c==1_180_591_620_717_411_303_433) x q[0]; auto analyzed = oq3::frontend::analyzeOpenQASM(source); ASSERT_TRUE(analyzed) << analyzed.diagnostics.front().message; EXPECT_TRUE(llvm::any_of(analyzed.program->conditions, [](const auto& c) { - return c.kind == oq3::frontend::ConditionKind::Bit && c.bit.index == 70; + return c.kind == oq3::frontend::ConditionKind::RegisterComparison && + c.expected[70]; })); } @@ -1711,12 +1727,10 @@ if(c==1) x q[0]; )qasm"; auto analyzed = oq3::frontend::analyzeOpenQASM(source); ASSERT_TRUE(analyzed) << analyzed.diagnostics.front().message; - // Truncating to 64 bits would omit Not(c[79]). - EXPECT_TRUE(llvm::any_of(analyzed.program->conditions, [&](const auto& c) { - return c.kind == oq3::frontend::ConditionKind::Not && - analyzed.program->conditions[c.lhs].kind == - oq3::frontend::ConditionKind::Bit && - analyzed.program->conditions[c.lhs].bit.index == 79; + /// Truncating to 64 bits would omit the leading zero bits. + EXPECT_TRUE(llvm::any_of(analyzed.program->conditions, [](const auto& c) { + return c.kind == oq3::frontend::ConditionKind::RegisterComparison && + c.expected.getBitWidth() == 80U && c.expected == 1U; })); } diff --git a/test/python/test_mlir_qiskit_translation.py b/test/python/test_mlir_qiskit_translation.py index 4cbf646df6..4eea8934c2 100644 --- a/test/python/test_mlir_qiskit_translation.py +++ b/test/python/test_mlir_qiskit_translation.py @@ -575,7 +575,98 @@ def test_cleanup_forwards_measurement_results_to_qiskit_condition() -> None: assert restored.data[2].operation.blocks[0].count_ops() == {"x": 1} condition = restored.data[2].operation.condition assert isinstance(condition, expr.Expr) - assert expr.structurally_equivalent(condition, expr.logic_and(*restored.clbits)) + assert expr.structurally_equivalent(condition, expr.equal(restored.cregs[0], 3)) + + +def test_openqasm_register_ordering_exports_to_qiskit_expression() -> None: + """Export first-class register ordering as a Qiskit Uint expression.""" + program = QCProgram.from_qasm_str( + """OPENQASM 3.1; +include "stdgates.inc"; +qubit[3] q; +bit[2] c; +c[0] = measure q[0]; +c[1] = measure q[1]; +if (c >= 1) { x q[2]; } +""" + ) + + restored = program.to_qiskit() + condition = restored.data[2].operation.condition + + assert "cbit.cmp uge" in program.ir + assert isinstance(condition, expr.Expr) + assert expr.structurally_equivalent(condition, expr.greater_equal(restored.cregs[0], 1)) + + reimported = QCProgram.from_qiskit(restored) + assert "cbit.cmp uge" in reimported.ir + reimported_circuit = reimported.to_qiskit() + reimported_condition = reimported_circuit.data[2].operation.condition + assert isinstance(reimported_condition, expr.Expr) + assert expr.structurally_equivalent(reimported_condition, expr.greater_equal(reimported_circuit.cregs[0], 1)) + + +@pytest.mark.parametrize( + ("comparison", "predicate"), + [ + (None, "eq"), + ("equal", "eq"), + ("not_equal", "ne"), + ("less", "ult"), + ("less_equal", "ule"), + ("greater", "ugt"), + ("greater_equal", "uge"), + ], +) +def test_qiskit_register_conditions_import_canonically(comparison: str | None, predicate: str) -> None: + """Import tuple and typed register conditions as first-class comparisons.""" + circuit = QuantumCircuit(1, 3) + condition = (circuit.cregs[0], 1) if comparison is None else getattr(expr, comparison)(circuit.cregs[0], 1) + with circuit.if_test(condition): + circuit.x(0) + + ir = QCProgram.from_qiskit(circuit).ir + + assert f"cbit.cmp {predicate}" in ir + assert "cbit.load" not in ir + + +@pytest.mark.parametrize( + ("comparison", "predicate"), + [ + ("equal", "eq"), + ("not_equal", "ne"), + ("less", "ugt"), + ("less_equal", "uge"), + ("greater", "ult"), + ("greater_equal", "ule"), + ], +) +def test_qiskit_reversed_register_conditions_import_canonically(comparison: str, predicate: str) -> None: + """Canonicalize Qiskit comparisons with the constant on the left.""" + circuit = QuantumCircuit(1, 3) + with circuit.if_test(getattr(expr, comparison)(1, circuit.cregs[0])): + circuit.x(0) + + ir = QCProgram.from_qiskit(circuit).ir + + assert f"cbit.cmp {predicate}" in ir + assert "cbit.load" not in ir + + +def test_qiskit_oversized_tuple_condition_is_false() -> None: + """Fold an impossible Qiskit register equality instead of widening it.""" + circuit = QuantumCircuit(1, 2) + with circuit.if_test((circuit.cregs[0], 4)): + circuit.x(0) + + program = QCProgram.from_qiskit(circuit) + + assert "arith.constant false" in program.ir + assert "cbit.cmp" not in program.ir + condition = program.to_qiskit().data[0].operation.condition + assert isinstance(condition, expr.Value) + assert condition.value == 0 def test_openqasm_short_circuit_expression_exports_to_qiskit() -> None: @@ -1906,12 +1997,12 @@ def test_uint_register_cast_to_bool_tests_all_bits() -> None: program = QCProgram.from_qiskit(circuit) ir = program.ir - assert "arith.cmpi ne" in ir + assert "cbit.cmp ne" in ir assert "arith.trunci" not in ir restored = program.to_qiskit() round_trip_ir = QCProgram.from_qiskit(restored).ir - assert "arith.cmpi ne" in round_trip_ir + assert "cbit.cmp ne" in round_trip_ir assert "arith.trunci" not in round_trip_ir @@ -2103,7 +2194,7 @@ def test_nested_condition_only_expression_uses_parent_capture_map() -> None: assert ir.count("scf.if") == 2 -def test_nested_legacy_clbit_condition_uses_root_index() -> None: +def test_nested_tuple_clbit_condition_uses_root_index() -> None: """Resolve a nested tuple condition through its enclosing Clbit map.""" circuit = QuantumCircuit(2, 2) with circuit.for_loop(range(2), None, None, None, None, label=None) as iteration: