From a56cf6c3acf4cb2d0bc43030fcef15339e9ff5f7 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Wed, 19 Aug 2026 17:17:37 +0200 Subject: [PATCH 01/12] =?UTF-8?q?=E2=9C=A8=20Extend=20the=20QCO=20DD=20cla?= =?UTF-8?q?ssical=20interpreter?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- CHANGELOG.md | 11 +- .../mlir/Dialect/QCO/Utils/DDFunctionality.h | 64 +- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 1303 ++++++++++++++--- .../QCO/Utils/test_dd_functionality.cpp | 592 +++++++- 4 files changed, 1670 insertions(+), 300 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index de2bda2db8..252867bf49 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,10 +33,12 @@ releases may include breaking changes. C++, Python, and command-line interfaces ([#2135]) ([**@denialhaag**], [**@burgholzer**]) - ✨ Add decision diagram-based construction, simulation, and sampling for QCO - programs, including mid-circuit `measure` / `reset`, concrete `if` / - `index_switch` / `scf.for` / `func.call`, classical SSA and CBit registers, - dense multi-wire embedding, output-aware multi-shot sampling, and Python - bindings ([#1915], [#1973], [#2077]) ([**@simon1hofmann**]) + programs, including mid-circuit `measure` / `reset`, concrete QCO and SCF + control flow, non-recursive calls, bound parameters, classical integer and + `f64` SSA, CBit registers, one-dimensional memrefs, dynamic quantum allocation + and qtensors, dense multi-wire embedding, output-aware multi-shot sampling, + and Python bindings ([#1915], [#1973], [#2077], [#2078]) + ([**@simon1hofmann**]) - ✨ Add immutable MLIR compiler targets, QDMI device integration, and target compilation through C++, Python, and `mqt-cc` ([#1687], [#1993], [#1999], [#2049]) ([**@MatthiasReumann**], [**@simon1hofmann**], [**@burgholzer**]) @@ -921,6 +923,7 @@ for previous changelogs._ [#2105]: https://github.com/munich-quantum-toolkit/core/pull/2105 [#2084]: https://github.com/munich-quantum-toolkit/core/pull/2084 [#2082]: https://github.com/munich-quantum-toolkit/core/pull/2082 +[#2078]: https://github.com/munich-quantum-toolkit/core/pull/2078 [#2077]: https://github.com/munich-quantum-toolkit/core/pull/2077 [#2074]: https://github.com/munich-quantum-toolkit/core/pull/2074 [#2066]: https://github.com/munich-quantum-toolkit/core/pull/2066 diff --git a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h index f3506f845a..bb4fdf22cf 100644 --- a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h +++ b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h @@ -12,7 +12,10 @@ #include "dd/Package_fwd.hpp" +#include #include +#include +#include #include #include @@ -22,6 +25,15 @@ namespace mlir::qco { +/** + * @brief Concrete values for symbolic QCO DD inputs. + * + * Integer and `f64` attributes bind scalar function arguments. An integer + * attribute bound to a dynamic one-dimensional qtensor argument gives its + * runtime extent. Bindings for other values are rejected. + */ +using DDBindings = DenseMap; + /** * @brief Sequentially build a matrix DD for a static unitary QCO `func.func`. * @@ -31,30 +43,36 @@ namespace mlir::qco { * wires, and applies unitary operations via decision-diagram multiplication. * * Supported programs: - * - Standard single-, two-, and three-qubit gates with compile-time constant + * - Standard single-, two-, and three-qubit gates with constant or bound * parameters (sparse DD path) * - `ctrl` with a sole standard-gate body (same sparse path) * - Other `UnitaryOpInterface` ops with a compile-time known matrix (`inv`, * compound `ctrl`, ...), including `gphase` and `barrier` + * - QTensor bookkeeping over existing input wires + * - Concrete QCO and SCF control flow and non-recursive single-block calls + * - Concrete integer, index, and `f64` arithmetic and one-dimensional memrefs + * of those scalar types * - `qco.static` establishes the wire map (or qubit-typed `func` args if none), - * followed by entry-block `qco.alloc`; `sink` is ignored; `arith.constant` - * is ignored for matrix construction; `func.return` accepts qubit results - * only in canonical wire order + * followed by entry-block `qco.alloc`; `sink` is ignored; returned qubits + * and qtensors must preserve canonical wire order * * Known one-, two-, and three-qubit matrices are constructed directly as DD * gates. Larger compile-time unitaries are embedded directly into a DD over * their target wires, so idle register qubits do not enlarge the local matrix. - * Measurements, resets, symbolic parameters, and control-flow ops are not - * supported. + * Measurements, resets, unbound parameters, and non-concrete control flow are + * not supported. * * @pre The containing module has passed MLIR verification and * `qco::verifyLinearity`. * * @param func The QCO function to construct the functionality for * @param dd The DD package to use (must hold at least the function's qubits) + * @param bindings Concrete values for symbolic function arguments * @return The matrix DD on success, or failure for unsupported programs */ -FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd); +FailureOr +buildFunctionality(func::FuncOp func, dd::Package& dd, + const DDBindings& bindings = DDBindings()); /** * @brief Simulate a QCO `func.func` that may contain measurements, resets, and @@ -63,17 +81,15 @@ FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd); * @details Supports the unitary op set of @ref buildFunctionality, plus * `qco.measure` / `qco.reset` (collapsing via @p rng) and `qco.if` / * `qco.index_switch` when the branch selector is a concrete classical SSA value - * (`arith.constant` integer/index, a prior measurement, a `cbit.load`, - * `arith.extui`, `arith.index_castui`, `arith.cmpi`, `arith.select`, - * `arith.addi` / `subi` / `muli`, or `andi` / `ori` / `xori` / `shli` / - * `shrui` on those values). The simulation tracks CBit initialization, loads, - * and stores. Only qubit-typed linear values are supported (no qtensors). - * Nested regions are walked; direct `scf.for` execution with concrete positive - * steps and non-recursive single-block `func.call` are supported. A shared - * 10000-step budget bounds loop iterations across nested loops and calls; - * `scf.while` and multi-block function bodies remain unsupported. - * Consumes one reference to @p in regardless of whether simulation succeeds or - * fails. + * (`arith.constant`, a prior measurement, integer and `f64` + * arithmetic, comparisons, casts, shifts, and `arith.select`). Dynamic quantum + * allocation, qtensors, memrefs, CBit registers, loops, regions, and calls are + * supported. QTensor sizes and indices must be concrete; dynamic qtensor + * arguments require an extent in @p bindings. A shared 10000-step budget bounds + * loop iterations across nested loops and calls. Multi-block function bodies + * remain unsupported. Allocated wires stay in the returned state after + * deallocation. + * Consumes one reference to @p in regardless of success or failure. * * @pre The containing module has passed MLIR verification and * `qco::verifyLinearity`. @@ -83,11 +99,13 @@ FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd); * higher wires are preserved; one reference is consumed * @param dd The DD package to use * @param rng RNG used for collapsing measurements and resets + * @param bindings Concrete values for symbolic function arguments * @return The output statevector DD on success, or failure for unsupported * programs */ FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, - dd::Package& dd, std::mt19937_64& rng); + dd::Package& dd, std::mt19937_64& rng, + const DDBindings& bindings = DDBindings()); /** * @brief Sample measurement outcomes from a QCO `func.func`. @@ -99,7 +117,8 @@ FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, * basis sampling via `Package::measureAll` (qubit `n-1` … `0`). Terminal entry- * block measurements that only produce returned CBit cells are sampled from * one DD evolution; resets and execution-dependent measurements are executed - * once per shot. + * once per shot. Dynamically allocated wires are included in fallback basis + * outcomes even after deallocation. * * @pre The containing module has passed MLIR verification and * `qco::verifyLinearity`. @@ -108,10 +127,11 @@ FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, * @param dd The DD package to use * @param shots Number of shots * @param rng RNG for collapsing measurements and non-collapsing sampling + * @param bindings Concrete values for symbolic function arguments * @return Histogram of outcome strings on success, or failure for unsupported * programs */ FailureOr> -sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng); - +sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng, + const DDBindings& bindings = DDBindings()); } // namespace mlir::qco diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index a9e52d540e..318f11f73e 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -29,14 +29,20 @@ #include "mlir/Dialect/QCO/IR/QCOInterfaces.h" #include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QCO/Utils/Matrix.h" +#include "mlir/Dialect/QTensor/IR/QTensorOps.h" +#include #include +#include #include +#include #include #include #include +#include #include #include +#include #include #include #include @@ -59,11 +65,14 @@ #include #include #include +#include #include namespace mlir::qco { namespace { +constexpr size_t MAX_CONTROL_FLOW_STEPS = 10'000; + struct QubitMap { DenseMap qubits; size_t numQubits = 0; @@ -107,17 +116,48 @@ struct QubitMap { } }; +/// Physical wires stored at each tensor index; extracted positions are empty. +using TensorSlots = SmallVector>; + +struct TensorMap { + DenseMap tensors; + + void bind(Value value, TensorSlots slots) { + tensors[value] = std::move(slots); + } + + [[nodiscard]] const TensorSlots* lookup(Value value) const { + const auto it = tensors.find(value); + return it == tensors.end() ? nullptr : &it->second; + } + + void erase(Value value) { tensors.erase(value); } +}; + struct ClassicalEnv { struct RegisterBit { std::optional value; std::optional deferredWire; }; using RegisterState = std::vector; + using Scalar = std::variant; - DenseMap scalars; + DenseMap values; DenseMap deferredMeasurements; /// Shared storage preserves CBit register identity across `func.call`. DenseMap> registers; + /// Shared storage preserves caller-visible writes through `func.call`. + DenseMap>> memrefs; + + LogicalResult bindFrom(Value source, Value dest, Operation* op) { + const auto it = values.find(source); + if (it == values.end()) { + return op->emitError() + << "classical SSA value is not mapped for QCO DD simulation"; + } + values[dest] = it->second; + return success(); + } }; struct DecodedGate { @@ -127,11 +167,12 @@ struct DecodedGate { struct WalkState { QubitMap* qubits; + TensorMap* tensors; ClassicalEnv* classical; dd::Package* dd; std::mt19937_64* rng = nullptr; const DenseSet* deferredMeasurements = nullptr; - size_t remainingExecutionSteps = 10'000; + size_t remainingExecutionSteps = MAX_CONTROL_FLOW_STEPS; DenseSet activeCalls; }; struct LoopRange { @@ -145,10 +186,41 @@ struct SamplingPlan { }; } // namespace +[[nodiscard]] static bool isQTensorType(Type type) { + const auto tensorType = dyn_cast(type); + return tensorType && tensorType.getRank() == 1 && + isa(tensorType.getElementType()); +} + +template +static FailureOr lookupScalar(Value value, const ClassicalEnv& classical, + Operation* op) { + const auto it = classical.values.find(value); + if (it == classical.values.end() || !std::holds_alternative(it->second)) { + return op->emitError() + << "classical SSA value is not mapped for QCO DD simulation"; + } + return std::get(it->second); +} + +static FailureOr +resolveDouble(Value value, const ClassicalEnv& classical, Operation* op) { + if (const auto it = classical.values.find(value); + it != classical.values.end() && + std::holds_alternative(it->second)) { + return std::get(it->second); + } + if (const auto constant = mqt::valueToDouble(value)) { + return *constant; + } + return op->emitError() + << "floating-point SSA value has no concrete QCO DD binding"; +} + /// `std::nullopt` if @p unitary is not a standard gate; failure if its unitary -/// matrix is not known at compile time. +/// parameters are not concrete. static FailureOr> -decodeStandardGate(UnitaryOpInterface unitary) { +decodeStandardGate(UnitaryOpInterface unitary, const ClassicalEnv& classical) { Operation* op = unitary.getOperation(); const auto type = TypeSwitch(op) @@ -185,15 +257,13 @@ decodeStandardGate(UnitaryOpInterface unitary) { if (type == qc::OpType::None) { return std::optional{std::nullopt}; } - if (!unitary.hasCompileTimeKnownUnitaryMatrix()) { - return unitary.emitError() - << "unitary must have a compile-time constant matrix"; - } - DecodedGate decoded{.type = type, .params = {}}; for (Value param : unitary.getParameters()) { - decoded.params.push_back( - static_cast(*mlir::mqt::valueToDouble(param))); + auto concrete = resolveDouble(param, classical, op); + if (failed(concrete)) { + return failure(); + } + decoded.params.push_back(static_cast(*concrete)); } return std::optional{std::move(decoded)}; } @@ -250,20 +320,23 @@ template static LogicalResult applyUnitaryMatrix(UnitaryOpInterface unitary, WalkState& walk, StateDD& state) { Operation* op = unitary.getOperation(); - if (!unitary.hasCompileTimeKnownUnitaryMatrix()) { - return unitary.emitError() - << "unitary must have a compile-time constant matrix"; - } if (auto gphase = dyn_cast(op)) { - const auto theta = *mlir::mqt::valueToDouble(gphase.getTheta()); + auto theta = resolveDouble(gphase.getTheta(), *walk.classical, op); + if (failed(theta)) { + return failure(); + } auto id = dd::Package::makeIdent(); - id.w = walk.dd->cn.lookup(std::cos(theta), std::sin(theta)); + id.w = walk.dd->cn.lookup(std::cos(*theta), std::sin(*theta)); state = walk.dd->applyOperation(id, state); return success(); } if (isa(op)) { return walk.qubits->remapUnitary(unitary); } + if (!unitary.hasCompileTimeKnownUnitaryMatrix()) { + return unitary.emitError() + << "unitary must have a compile-time constant matrix"; + } DynamicMatrix local; if (!unitary.getUnitaryMatrixDynamic(local)) { @@ -343,9 +416,25 @@ static LogicalResult applyDecodedStandard(UnitaryOpInterface unitary, } static LogicalResult validateReturn(func::ReturnOp returnOp, - const QubitMap& qubits) { + const QubitMap& qubits, + const TensorMap& tensors) { qc::Qubit expected = 0; for (Value value : returnOp.getOperands()) { + if (isQTensorType(value.getType())) { + const auto* slots = tensors.lookup(value); + if (slots == nullptr) { + return returnOp.emitError() + << "returned qtensor is not mapped for QCO DD simulation"; + } + for (const auto wire : *slots) { + if (!wire || *wire != expected) { + return returnOp.emitError() + << "returned qubits must preserve canonical wire order"; + } + ++expected; + } + continue; + } if (!isa(value.getType())) { continue; } @@ -367,49 +456,152 @@ static LogicalResult validateReturn(func::ReturnOp returnOp, return success(); } -static void bindInteger(Value result, const llvm::APInt& value, - ClassicalEnv& classical) { - classical.scalars[result] = IntegerAttr::get(result.getType(), value); +static LogicalResult recordConstant(arith::ConstantOp constant, + ClassicalEnv& classical) { + if (auto attr = dyn_cast(constant.getValue())) { + classical.values[constant.getResult()] = attr.getValue(); + } else if (auto attr = dyn_cast(constant.getValue())) { + if (!constant.getType().isF64()) { + return constant.emitError() + << "QCO DD simulation only supports f64 classical values"; + } + classical.values[constant.getResult()] = attr.getValue().convertToDouble(); + } else if (auto attr = dyn_cast(constant.getValue())) { + if (constant.getType().isInteger(1)) { + classical.values[constant.getResult()] = attr.getValue() != 0; + } else if (isa(constant.getType())) { + classical.values[constant.getResult()] = attr.getInt(); + } else if (isa(constant.getType())) { + classical.values[constant.getResult()] = attr.getValue(); + } + } + return success(); } -static FailureOr -lookupInteger(Value value, ClassicalEnv& classical, Operation* op) { - const auto it = classical.scalars.find(value); - if (it == classical.scalars.end()) { - return op->emitError() << "classical SSA value is not mapped for QCO DD " - "simulation: " - << value.getType(); - } - return it->second.getValue(); +static LogicalResult applyBindings(func::FuncOp func, + const DDBindings& bindings, + ClassicalEnv& classical) { + for (const auto& [value, attr] : bindings) { + auto argument = dyn_cast(value); + if (!argument || argument.getOwner() != &func.getBody().front()) { + return func.emitError() + << "QCO DD bindings must target entry-block arguments"; + } + const Type type = value.getType(); + if (isQTensorType(type)) { + if (cast(type).isDynamicDim(0) && + isa(attr)) { + continue; + } + } else if (type.isInteger(1)) { + if (auto boolean = dyn_cast(attr)) { + classical.values[value] = boolean.getValue(); + continue; + } + if (auto integer = dyn_cast(attr)) { + classical.values[value] = integer.getValue() != 0; + continue; + } + } else if (isa(type)) { + if (auto integer = dyn_cast(attr)) { + classical.values[value] = integer.getInt(); + continue; + } + } else if (auto integerType = dyn_cast(type)) { + if (auto integer = dyn_cast(attr)) { + classical.values[value] = + integer.getValue().sextOrTrunc(integerType.getWidth()); + continue; + } + } else if (type.isF64()) { + if (auto floating = dyn_cast(attr); + floating && floating.getType().isF64()) { + classical.values[value] = floating.getValue().convertToDouble(); + continue; + } + } + return func.emitError() << "QCO DD binding attribute " << attr + << " does not match argument type " << type; + } + return success(); } -static FailureOr lookupBool(Value value, ClassicalEnv& classical, +static FailureOr lookupBool(Value value, const ClassicalEnv& classical, Operation* op) { - auto result = lookupInteger(value, classical, op); - if (failed(result)) { - return failure(); - } - return !result->isZero(); + return lookupScalar(value, classical, op); } -static FailureOr lookupIndex(Value value, ClassicalEnv& classical, - Operation* op) { - auto result = lookupInteger(value, classical, op); - if (failed(result)) { - return failure(); +static FailureOr +lookupIndex(Value value, const ClassicalEnv& classical, Operation* op) { + return lookupScalar(value, classical, op); +} + +static FailureOr lookupFloat(Value value, const ClassicalEnv& classical, + Operation* op) { + return lookupScalar(value, classical, op); +} + +static FailureOr +lookupInteger(Value value, const ClassicalEnv& classical, Operation* op) { + if (value.getType().isInteger(1)) { + auto bit = lookupBool(value, classical, op); + if (failed(bit)) { + return failure(); + } + return llvm::APInt(1, static_cast(*bit)); + } + if (isa(value.getType())) { + auto index = lookupIndex(value, classical, op); + if (failed(index)) { + return failure(); + } + return llvm::APInt(64, static_cast(*index)); } - return result->getSExtValue(); + if (isa(value.getType())) { + return lookupScalar(value, classical, op); + } + return op->emitError() << "expected an integer or index SSA value"; } -static LogicalResult applyUnsignedIndexCast(Value in, Value out, Operation* op, - ClassicalEnv& classical) { - auto value = lookupInteger(in, classical, op); - if (failed(value)) { +[[nodiscard]] static bool evaluateCmp(arith::CmpIPredicate predicate, + const llvm::APInt& lhs, + const llvm::APInt& rhs) { + switch (predicate) { + case arith::CmpIPredicate::eq: + return lhs == rhs; + case arith::CmpIPredicate::ne: + return lhs != rhs; + case arith::CmpIPredicate::slt: + return lhs.slt(rhs); + case arith::CmpIPredicate::sle: + return lhs.sle(rhs); + case arith::CmpIPredicate::sgt: + return lhs.sgt(rhs); + case arith::CmpIPredicate::sge: + return lhs.sge(rhs); + case arith::CmpIPredicate::ult: + return lhs.ult(rhs); + case arith::CmpIPredicate::ule: + return lhs.ule(rhs); + case arith::CmpIPredicate::ugt: + return lhs.ugt(rhs); + case arith::CmpIPredicate::uge: + return lhs.uge(rhs); + } + llvm_unreachable("unknown arith.cmpi predicate"); +} + +static LogicalResult bindInteger(Value dest, const llvm::APInt& value, + ClassicalEnv& classical) { + if (dest.getType().isInteger(1)) { + classical.values[dest] = value[0]; + } else if (isa(dest.getType())) { + classical.values[dest] = static_cast(value.getZExtValue()); + } else if (auto type = dyn_cast(dest.getType())) { + classical.values[dest] = value.zextOrTrunc(type.getWidth()); + } else { return failure(); } - const auto integerType = dyn_cast(out.getType()); - const unsigned width = integerType ? integerType.getWidth() : 64U; - bindInteger(out, value->zextOrTrunc(width), classical); return success(); } @@ -428,7 +620,7 @@ static LogicalResult allocateRegister(cbit::AllocOp alloc, static FailureOr resolveRegisterIndex(Value index, cbit::RegisterType type, - ClassicalEnv& classical, + const ClassicalEnv& classical, Operation* op) { auto resolved = lookupIndex(index, classical, op); if (failed(resolved)) { @@ -485,86 +677,390 @@ static LogicalResult loadRegister(cbit::LoadOp load, ClassicalEnv& classical) { if (!cell.value) { return load.emitError() << "read from an undefined CBit register element"; } - bindInteger(load.getResult(), llvm::APInt(1, *cell.value ? 1 : 0), classical); + return bindInteger(load.getResult(), + llvm::APInt(1, static_cast(*cell.value)), + classical); +} + +[[nodiscard]] static bool isSupportedClassicalType(Type type) { + return isa(type) || type.isF64(); +} + +static FailureOr +lookupMemRefSlot(Value memref, ValueRange indices, ClassicalEnv& classical, + Operation* op) { + const auto type = dyn_cast(memref.getType()); + if (!type || type.getRank() != 1 || indices.size() != 1 || + !isSupportedClassicalType(type.getElementType())) { + return op->emitError() + << "QCO DD simulation only supports one-dimensional memrefs of " + "integer, index, or f64 values"; + } + auto index = lookupIndex(indices[0], classical, op); + if (failed(index)) { + return failure(); + } + const auto it = classical.memrefs.find(memref); + if (it == classical.memrefs.end()) { + return op->emitError() + << "classical memref is not mapped for QCO DD simulation"; + } + if (*index < 0 || static_cast(*index) >= it->second->size()) { + return op->emitError() + << "classical memref index out of range for QCO DD simulation"; + } + return &(*it->second)[static_cast(*index)]; +} + +static ClassicalEnv::Scalar zeroScalar(Type type) { + if (type.isInteger(1)) { + return false; + } + if (isa(type)) { + return int64_t{0}; + } + if (auto integer = dyn_cast(type)) { + return llvm::APInt(integer.getWidth(), 0); + } + return 0.0; +} + +static LogicalResult applyMemRefAlloc(memref::AllocOp alloc, + ClassicalEnv& classical) { + const auto type = dyn_cast(alloc.getType()); + if (!type || type.getRank() != 1 || + !isSupportedClassicalType(type.getElementType())) { + return alloc.emitError() + << "QCO DD simulation only supports one-dimensional memrefs of " + "integer, index, or f64 values"; + } + if (!alloc.getSymbolOperands().empty()) { + return alloc.emitError() + << "QCO DD simulation does not support symbolic memref operands"; + } + int64_t size = type.getDimSize(0); + if (type.isDynamicDim(0)) { + if (alloc.getDynamicSizes().size() != 1) { + return alloc.emitError() << "dynamic 1-D memref requires one size"; + } + auto dynamicSize = + lookupIndex(alloc.getDynamicSizes()[0], classical, alloc); + if (failed(dynamicSize)) { + return failure(); + } + size = *dynamicSize; + } + if (size < 0) { + return alloc.emitError() << "classical memref size must be non-negative"; + } + classical.memrefs[alloc.getResult()] = + std::make_shared>( + static_cast(size), zeroScalar(type.getElementType())); return success(); } -static LogicalResult applyBinaryInteger(Operation& op, - ClassicalEnv& classical) { - auto lhs = lookupInteger(op.getOperand(0), classical, &op); - auto rhs = lookupInteger(op.getOperand(1), classical, &op); +static LogicalResult applyMemRefStore(memref::StoreOp store, + ClassicalEnv& classical) { + auto slot = + lookupMemRefSlot(store.getMemref(), store.getIndices(), classical, store); + const auto value = classical.values.find(store.getValue()); + if (failed(slot) || value == classical.values.end()) { + if (value == classical.values.end()) { + store.emitError() + << "stored classical value is not mapped for QCO DD simulation"; + } + return failure(); + } + **slot = value->second; + return success(); +} + +static LogicalResult applyMemRefLoad(memref::LoadOp load, + ClassicalEnv& classical) { + auto slot = + lookupMemRefSlot(load.getMemref(), load.getIndices(), classical, load); + if (failed(slot)) { + return failure(); + } + classical.values[load.getResult()] = **slot; + return success(); +} + +template +static LogicalResult applyBinaryInteger(OpTy op, ClassicalEnv& classical, + Combine combine) { + auto lhs = lookupInteger(op.getLhs(), classical, op); + auto rhs = lookupInteger(op.getRhs(), classical, op); if (failed(lhs) || failed(rhs)) { return failure(); } - llvm::APInt result = *lhs; - if (isa(&op)) { - result &= *rhs; - } else if (isa(&op)) { - result |= *rhs; - } else if (isa(&op)) { - result ^= *rhs; - } else if (isa(&op)) { - result += *rhs; - } else if (isa(&op)) { - result -= *rhs; - } else if (isa(&op)) { - result *= *rhs; - } else { - if (rhs->isNegative() || rhs->uge(lhs->getBitWidth())) { - return op.emitError() - << "shift amount out of range for QCO DD simulation"; - } - const auto amount = static_cast(rhs->getZExtValue()); - result = isa(&op) ? lhs->shl(amount) : lhs->lshr(amount); + return bindInteger(op.getResult(), combine(*lhs, *rhs), classical); +} + +template +static LogicalResult applyBinaryFloat(OpTy op, ClassicalEnv& classical, + Combine combine) { + auto lhs = lookupFloat(op.getLhs(), classical, op); + auto rhs = lookupFloat(op.getRhs(), classical, op); + if (failed(lhs) || failed(rhs)) { + return failure(); } - bindInteger(op.getResult(0), result, classical); + classical.values[op.getResult()] = combine(*lhs, *rhs); return success(); } +template +static LogicalResult applyDivision(OpTy op, ClassicalEnv& classical, + Combine combine) { + auto rhs = lookupInteger(op.getRhs(), classical, op); + if (failed(rhs)) { + return failure(); + } + if (rhs->isZero()) { + return op.emitError() << "division by zero during QCO DD simulation"; + } + auto lhs = lookupInteger(op.getLhs(), classical, op); + if (failed(lhs)) { + return failure(); + } + return bindInteger(op.getResult(), combine(*lhs, *rhs), classical); +} + +template +static LogicalResult applyShift(OpTy op, ClassicalEnv& classical, Shift shift) { + auto lhs = lookupInteger(op.getLhs(), classical, op); + auto rhs = lookupInteger(op.getRhs(), classical, op); + if (failed(lhs) || failed(rhs)) { + return failure(); + } + if (rhs->uge(lhs->getBitWidth())) { + return op.emitError() << "shift amount out of range for QCO DD simulation"; + } + return bindInteger(op.getResult(), shift(*lhs, rhs->getZExtValue()), + classical); +} + +static LogicalResult applyIntegerCast(Value in, Value out, Operation* op, + ClassicalEnv& classical, bool isSigned) { + auto value = lookupInteger(in, classical, op); + if (failed(value)) { + return failure(); + } + const unsigned width = isa(out.getType()) + ? 64U + : cast(out.getType()).getWidth(); + if (width > value->getBitWidth()) { + *value = isSigned ? value->sext(width) : value->zext(width); + } else if (width < value->getBitWidth()) { + *value = value->trunc(width); + } + return bindInteger(out, *value, classical); +} + static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { + const auto isUnsupportedFloat = [](Type type) { + return isa(type) && !type.isF64(); + }; + if (llvm::any_of(op.getOperandTypes(), isUnsupportedFloat) || + llvm::any_of(op.getResultTypes(), isUnsupportedFloat)) { + return op.emitError() + << "QCO DD simulation only supports f64 classical values"; + } return TypeSwitch(&op) - .Case( - [&](Operation* binary) { - return applyBinaryInteger(*binary, classical); - }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs & rhs; + }); + }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs | rhs; + }); + }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs ^ rhs; + }); + }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs + rhs; + }); + }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs - rhs; + }); + }) + .Case([&](auto value) { + return applyBinaryInteger( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs * rhs; + }); + }) + .Case([&](auto value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.udiv(rhs); + }); + }) + .Case([&](auto value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.sdiv(rhs); + }); + }) + .Case([&](auto value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.urem(rhs); + }); + }) + .Case([&](auto value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.srem(rhs); + }); + }) + .Case([&](auto value) { + return applyShift( + value, classical, + [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.shl(rhs); }); + }) + .Case([&](auto value) { + return applyShift( + value, classical, + [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.lshr(rhs); }); + }) + .Case([&](auto value) { + return applyShift( + value, classical, + [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.ashr(rhs); }); + }) .Case([&](arith::CmpIOp cmp) -> LogicalResult { auto lhs = lookupInteger(cmp.getLhs(), classical, cmp); auto rhs = lookupInteger(cmp.getRhs(), classical, cmp); if (failed(lhs) || failed(rhs)) { return failure(); } - bindInteger(cmp.getResult(), - llvm::APInt(1, arith::applyCmpPredicate(cmp.getPredicate(), - *lhs, *rhs)), - classical); + classical.values[cmp.getResult()] = + evaluateCmp(cmp.getPredicate(), *lhs, *rhs); return success(); }) .Case([&](arith::SelectOp select) -> LogicalResult { - auto cond = lookupBool(select.getCondition(), classical, select); - if (failed(cond)) { + auto condition = lookupBool(select.getCondition(), classical, select); + if (failed(condition)) { return failure(); } - if (!isa(select.getType())) { - return select.emitError() - << "QCO DD simulation only supports integer or index select"; - } - auto t = lookupInteger(select.getTrueValue(), classical, select); - auto f = lookupInteger(select.getFalseValue(), classical, select); - if (failed(t) || failed(f)) { + Value selected = + *condition ? select.getTrueValue() : select.getFalseValue(); + return classical.bindFrom(selected, select.getResult(), select); + }) + .Case([&](arith::ExtUIOp ext) { + return applyIntegerCast(ext.getIn(), ext.getOut(), ext, classical, + false); + }) + .Case([&](arith::ExtSIOp cast) { + return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, + true); + }) + .Case([&](arith::IndexCastUIOp cast) { + return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, + false); + }) + .Case([&](arith::IndexCastOp cast) { + return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, + true); + }) + .Case([&](arith::TruncIOp cast) { + return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, + false); + }) + .Case([&](auto value) { + return applyBinaryFloat( + value, classical, [](double lhs, double rhs) { return lhs + rhs; }); + }) + .Case([&](auto value) { + return applyBinaryFloat( + value, classical, [](double lhs, double rhs) { return lhs - rhs; }); + }) + .Case([&](auto value) { + return applyBinaryFloat( + value, classical, [](double lhs, double rhs) { return lhs * rhs; }); + }) + .Case([&](auto value) { + return applyBinaryFloat( + value, classical, [](double lhs, double rhs) { return lhs / rhs; }); + }) + .Case([&](auto value) { + return applyBinaryFloat(value, classical, [](double lhs, double rhs) { + return std::fmod(lhs, rhs); + }); + }) + .Case([&](arith::NegFOp neg) -> LogicalResult { + auto value = lookupFloat(neg.getOperand(), classical, neg); + if (failed(value)) { return failure(); } - bindInteger(select.getResult(), *cond ? *t : *f, classical); + classical.values[neg.getResult()] = -*value; return success(); }) - .Case([&](arith::IndexCastUIOp cast) { - return applyUnsignedIndexCast(cast.getIn(), cast.getOut(), cast, - classical); - }) - .Case([&](arith::ExtUIOp cast) { - return applyUnsignedIndexCast(cast.getIn(), cast.getOut(), cast, - classical); + .Case([&](arith::CmpFOp cmp) -> LogicalResult { + auto lhs = lookupFloat(cmp.getLhs(), classical, cmp); + auto rhs = lookupFloat(cmp.getRhs(), classical, cmp); + if (failed(lhs) || failed(rhs)) { + return failure(); + } + classical.values[cmp.getResult()] = arith::applyCmpPredicate( + cmp.getPredicate(), llvm::APFloat(*lhs), llvm::APFloat(*rhs)); + return success(); }) + .Case( + [&](Operation* castOp) -> LogicalResult { + auto value = + lookupInteger(castOp->getOperand(0), classical, castOp); + if (failed(value)) { + return failure(); + } + classical.values[castOp->getResult(0)] = + value->roundToDouble(isa(castOp)); + return success(); + }) + .Case( + [&](Operation* castOp) -> LogicalResult { + auto value = lookupFloat(castOp->getOperand(0), classical, castOp); + if (failed(value)) { + return failure(); + } + Value out = castOp->getResult(0); + const unsigned width = cast(out.getType()).getWidth(); + const bool isSigned = isa(castOp); + llvm::APSInt result(width, /*isUnsigned=*/!isSigned); + bool exact = false; + const auto status = llvm::APFloat(*value).convertToInteger( + result, llvm::APFloat::rmTowardZero, &exact); + if ((status & llvm::APFloat::opInvalidOp) != 0) { + return castOp->emitError() + << "floating-point value is outside the destination " + "integer range during QCO DD simulation"; + } + return bindInteger(out, result, classical); + }) .Default([](Operation* unsupported) { return unsupported->emitError() << "unsupported classical op for QCO DD simulation: " @@ -608,9 +1104,24 @@ resolveLoop(scf::ForOp forOp, ClassicalEnv& classical, size_t remainingSteps) { return LoopRange{.induction = lowerWide, .step = stepWide, .trips = limited}; } +static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, + WalkState& walk, Operation* op); + +static LogicalResult bindLinearArgs(ValueRange operands, Block& block, + WalkState& walk, Operation* op) { + for (Value arg : block.getArguments()) { + if (!isa(arg.getType()) && !isQTensorType(arg.getType())) { + return op->emitError() + << "unsupported linear region argument for QCO DD simulation"; + } + } + return bindValuePairs(operands, block.getArguments(), walk, op); +} + static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, WalkState& walk, Operation* op) { const QubitMap sourceQubits = *walk.qubits; + const TensorMap sourceTensors = *walk.tensors; const ClassicalEnv sourceClassical = *walk.classical; for (auto [src, dest] : llvm::zip_equal(sources, dests)) { if (isa(dest.getType())) { @@ -620,6 +1131,13 @@ static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, << "qubit SSA value is not mapped for QCO DD construction"; } walk.qubits->bind(dest, *q); + } else if (isQTensorType(dest.getType())) { + const auto* slots = sourceTensors.lookup(src); + if (slots == nullptr) { + return op->emitError() + << "qtensor SSA value is not mapped for QCO DD simulation"; + } + walk.tensors->bind(dest, *slots); } else if (isa(dest.getType())) { const auto it = sourceClassical.registers.find(src); if (it == sourceClassical.registers.end()) { @@ -627,18 +1145,42 @@ static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, << "CBit register is not mapped for QCO DD simulation"; } walk.classical->registers[dest] = it->second; + } else if (isa(dest.getType())) { + const auto it = sourceClassical.memrefs.find(src); + if (it == sourceClassical.memrefs.end()) { + return op->emitError() + << "classical memref is not mapped for QCO DD simulation"; + } + walk.classical->memrefs[dest] = it->second; } else { - const auto value = sourceClassical.scalars.find(src); - if (value == sourceClassical.scalars.end()) { + const auto value = sourceClassical.values.find(src); + if (value == sourceClassical.values.end()) { return op->emitError() << "classical SSA value is not mapped for QCO DD simulation"; } - walk.classical->scalars[dest] = value->second; + walk.classical->values[dest] = value->second; } } return success(); } +static LogicalResult bindYieldResults(YieldOp yield, + ValueRange classicalResults, + ValueRange linearResults, + WalkState& walk) { + const size_t numClassical = classicalResults.size(); + if (yield.getNumOperands() != numClassical + linearResults.size()) { + return yield.emitError() + << "yield operand count does not match operation results"; + } + if (failed(bindValuePairs(yield.getOperands().take_front(numClassical), + classicalResults, walk, yield))) { + return failure(); + } + return bindValuePairs(yield.getOperands().drop_front(numClassical), + linearResults, walk, yield); +} + template static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state); @@ -653,29 +1195,197 @@ static LogicalResult walkBlock(Block& block, WalkState& walk, StateDD& state) { } template -static LogicalResult applyRegionBranch(ValueRange linearOperands, Block& block, - WalkState& walk, StateDD& state, - Operation* parent) { - if (failed( - bindValuePairs(linearOperands, block.getArguments(), walk, parent))) { +static LogicalResult +applyRegionBranch(ValueRange linearOperands, Block& block, + ValueRange classicalResults, ValueRange linearResults, + WalkState& walk, StateDD& state, Operation* parent) { + if (failed(bindLinearArgs(linearOperands, block, walk, parent))) { return failure(); } if (failed(walkBlock(block, walk, state))) { return failure(); } - auto yield = cast(block.getTerminator()); - return bindValuePairs(yield.getOperands(), parent->getResults(), walk, yield); + return bindYieldResults(cast(block.getTerminator()), + classicalResults, linearResults, walk); +} + +template +static LogicalResult applyScfRegion(Region& region, ValueRange results, + WalkState& walk, StateDD& state, + Operation* parent) { + if (!region.hasOneBlock()) { + return parent->emitError() + << "SCF region must contain exactly one block for QCO DD simulation"; + } + Block& block = region.front(); + if (failed(walkBlock(block, walk, state))) { + return failure(); + } + auto yield = dyn_cast(block.getTerminator()); + if (!yield || yield.getNumOperands() != results.size()) { + return parent->emitError() + << "SCF region must yield one value for each result"; + } + return bindValuePairs(yield.getOperands(), results, walk, parent); +} + +static FailureOr allocateZeroQubits(size_t count, WalkState& walk, + dd::VectorDD& state, + Operation* op) { + if (count == 0) { + return op->emitError() << "quantum allocation size must be positive"; + } + if (walk.qubits->numQubits > walk.dd->qubits() || + count > walk.dd->qubits() - walk.qubits->numQubits) { + return op->emitError() << "DD package has " << walk.dd->qubits() + << " qubits but allocation requires " + << walk.qubits->numQubits + count; + } + + const size_t first = walk.qubits->numQubits; + auto zeros = dd::makeZeroState(count, *walk.dd, first); + auto extended = walk.dd->kronecker(zeros, state, first, /*incIdx=*/false); + walk.dd->incRef(extended); + walk.dd->decRef(zeros); + walk.dd->decRef(state); + state = extended; + + TensorSlots slots; + slots.reserve(count); + for (size_t i = 0; i < count; ++i) { + slots.emplace_back(static_cast(first + i)); + } + walk.qubits->numQubits += count; + return slots; } template static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return TypeSwitch(&op) - .template Case([](auto) { return success(); }) + .template Case([](auto) { return success(); }) .template Case([&](arith::ConstantOp constant) { - if (auto attr = dyn_cast(constant.getValue())) { - walk.classical->scalars[constant.getResult()] = attr; + return recordConstant(constant, *walk.classical); + }) + .template Case([&](AllocOp alloc) -> LogicalResult { + if constexpr (!std::is_same_v) { + if (!walk.qubits->lookup(alloc.getResult())) { + return alloc.emitError() + << "dynamic qubit allocation is not supported for QCO DD " + "functionality construction"; + } + return success(); + } else { + auto slots = allocateZeroQubits(1, walk, state, alloc); + if (failed(slots)) { + return failure(); + } + walk.qubits->bind(alloc.getResult(), *slots->front()); + return success(); } - return success(); + }) + .template Case( + [&](qtensor::AllocOp alloc) -> LogicalResult { + if constexpr (!std::is_same_v) { + return alloc.emitError() + << "qtensor allocation is not supported for QCO DD " + "functionality construction"; + } else { + auto size = lookupIndex(alloc.getSize(), *walk.classical, alloc); + if (failed(size)) { + return failure(); + } + if (*size <= 0) { + return alloc.emitError() + << "qtensor allocation size must be positive"; + } + auto slots = allocateZeroQubits(static_cast(*size), walk, + state, alloc); + if (failed(slots)) { + return failure(); + } + walk.tensors->bind(alloc.getResult(), std::move(*slots)); + return success(); + } + }) + .template Case( + [&](qtensor::FromElementsOp fromElements) -> LogicalResult { + auto wires = walk.qubits->lookupRange(fromElements.getElements(), + fromElements); + if (failed(wires)) { + return failure(); + } + TensorSlots slots; + slots.reserve(wires->size()); + for (const qc::Qubit wire : *wires) { + slots.emplace_back(wire); + } + walk.tensors->bind(fromElements.getResult(), std::move(slots)); + return success(); + }) + .template Case( + [&](qtensor::ExtractOp extract) -> LogicalResult { + const auto* input = walk.tensors->lookup(extract.getTensor()); + auto index = + lookupIndex(extract.getIndex(), *walk.classical, extract); + if (input == nullptr || failed(index)) { + if (input == nullptr) { + extract.emitError() + << "qtensor is not mapped for QCO DD simulation"; + } + return failure(); + } + if (*index < 0 || static_cast(*index) >= input->size()) { + return extract.emitError() << "qtensor index out of range"; + } + TensorSlots output = *input; + auto& wire = output[static_cast(*index)]; + if (!wire) { + return extract.emitError() + << "qtensor element has already been extracted"; + } + walk.qubits->bind(extract.getResult(), *wire); + wire.reset(); + walk.tensors->bind(extract.getOutTensor(), std::move(output)); + return success(); + }) + .template Case( + [&](qtensor::InsertOp insert) -> LogicalResult { + const auto* input = walk.tensors->lookup(insert.getDest()); + const auto wire = walk.qubits->lookup(insert.getScalar()); + auto index = + lookupIndex(insert.getIndex(), *walk.classical, insert); + if (input == nullptr || !wire || failed(index)) { + if (input == nullptr || !wire) { + insert.emitError() + << "qtensor or qubit is not mapped for QCO DD simulation"; + } + return failure(); + } + if (*index < 0 || static_cast(*index) >= input->size()) { + return insert.emitError() << "qtensor index out of range"; + } + TensorSlots output = *input; + output[static_cast(*index)] = wire; + walk.tensors->bind(insert.getResult(), std::move(output)); + return success(); + }) + .template Case( + [&](qtensor::DeallocOp dealloc) -> LogicalResult { + if (walk.tensors->lookup(dealloc.getTensor()) == nullptr) { + return dealloc.emitError() + << "qtensor is not mapped for QCO DD simulation"; + } + walk.tensors->erase(dealloc.getTensor()); + return success(); + }) + .template Case([&](memref::AllocOp alloc) { + return applyMemRefAlloc(alloc, *walk.classical); + }) + .template Case([&](memref::StoreOp store) { + return applyMemRefStore(store, *walk.classical); + }) + .template Case([&](memref::LoadOp load) { + return applyMemRefLoad(load, *walk.classical); }) .template Case([&](cbit::AllocOp alloc) { return allocateRegister(alloc, *walk.classical); @@ -686,15 +1396,21 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { .template Case([&](cbit::StoreOp store) { return storeRegister(store, *walk.classical); }) - .template Case( + .template Case([](auto) { return success(); }) + .template Case< + arith::AndIOp, arith::OrIOp, arith::XOrIOp, arith::AddIOp, + arith::SubIOp, arith::MulIOp, arith::DivUIOp, arith::DivSIOp, + arith::RemUIOp, arith::RemSIOp, arith::ShLIOp, arith::ShRUIOp, + arith::ShRSIOp, arith::CmpIOp, arith::SelectOp, arith::ExtUIOp, + arith::ExtSIOp, arith::IndexCastUIOp, arith::IndexCastOp, + arith::TruncIOp, arith::AddFOp, arith::SubFOp, arith::MulFOp, + arith::DivFOp, arith::RemFOp, arith::NegFOp, arith::CmpFOp, + arith::SIToFPOp, arith::UIToFPOp, arith::FPToSIOp, arith::FPToUIOp>( [&](Operation* classicalOp) { return applyClassicalOp(*classicalOp, *walk.classical); }) .template Case([&](func::ReturnOp returnOp) { - return validateReturn(returnOp, *walk.qubits); + return validateReturn(returnOp, *walk.qubits, *walk.tensors); }) .template Case([&](MeasureOp measureOp) -> LogicalResult { if constexpr (!std::is_same_v) { @@ -720,8 +1436,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return success(); } const char bit = walk.dd->measureOneCollapsing(state, *q, *walk.rng); - bindInteger(measureOp.getResult(), llvm::APInt(1, bit == '1'), - *walk.classical); + walk.classical->values[measureOp.getResult()] = bit == '1'; walk.qubits->bind(measureOp.getQubitOut(), *q); return success(); } @@ -752,84 +1467,163 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { } }) .template Case([&](IfOp ifOp) -> LogicalResult { - if constexpr (!std::is_same_v) { - return ifOp.emitError() - << "control-flow is not supported for QCO DD functionality " - "construction"; - } else { - auto condition = - lookupBool(ifOp.getCondition(), *walk.classical, ifOp); - if (failed(condition)) { - return failure(); - } - Block* block = *condition ? ifOp.thenBlock() : ifOp.elseBlock(); - return applyRegionBranch(ifOp.getQubits(), *block, walk, state, ifOp); + auto condition = lookupBool(ifOp.getCondition(), *walk.classical, ifOp); + if (failed(condition)) { + return failure(); } + Block* block = *condition ? ifOp.thenBlock() : ifOp.elseBlock(); + if (block == nullptr) { + return ifOp.emitError() << "selected qco.if region is empty"; + } + return applyRegionBranch(ifOp.getQubits(), *block, + ifOp.getClassicalResults(), + ifOp.getLinearResults(), walk, state, ifOp); }) .template Case( [&](IndexSwitchOp switchOp) -> LogicalResult { - if constexpr (!std::is_same_v) { - return switchOp.emitError() - << "control-flow is not supported for QCO DD " - "functionality construction"; - } else { - auto index = - lookupIndex(switchOp.getArg(), *walk.classical, switchOp); - if (failed(index)) { - return failure(); + auto selector = + lookupIndex(switchOp.getArg(), *walk.classical, switchOp); + if (failed(selector)) { + return failure(); + } + Block* block = switchOp.getDefaultBlock(); + for (auto [i, caseValue] : llvm::enumerate(switchOp.getCases())) { + if (caseValue == *selector) { + block = switchOp.getCaseBlock(i); + break; } - const auto cases = switchOp.getCases(); - Block* block = switchOp.getDefaultBlock(); - for (auto [i, caseValue] : llvm::enumerate(cases)) { - if (caseValue == *index) { - block = switchOp.getCaseBlock(i); - break; - } + } + if (block == nullptr) { + return switchOp.emitError() + << "selected qco.index_switch region is empty"; + } + return applyRegionBranch( + switchOp.getTargets(), *block, switchOp.getClassicalResults(), + switchOp.getLinearResults(), walk, state, switchOp); + }) + .template Case([&](scf::IfOp ifOp) -> LogicalResult { + auto condition = lookupBool(ifOp.getCondition(), *walk.classical, ifOp); + if (failed(condition)) { + return failure(); + } + Region& selected = + *condition ? ifOp.getThenRegion() : ifOp.getElseRegion(); + if (selected.empty()) { + return ifOp.getNumResults() == 0 + ? success() + : ifOp.emitError() + << "selected empty scf.if region has results"; + } + return applyScfRegion(selected, ifOp.getResults(), walk, state, ifOp); + }) + .template Case( + [&](scf::IndexSwitchOp switchOp) -> LogicalResult { + auto selector = + lookupIndex(switchOp.getArg(), *walk.classical, switchOp); + if (failed(selector)) { + return failure(); + } + Region* selected = &switchOp.getDefaultRegion(); + for (auto [i, value] : llvm::enumerate(switchOp.getCases())) { + if (value == *selector) { + selected = &switchOp.getCaseRegions()[i]; + break; } - return applyRegionBranch(switchOp.getTargets(), *block, walk, - state, switchOp); } + return applyScfRegion(*selected, switchOp.getResults(), walk, state, + switchOp); + }) + .template Case( + [&](scf::ExecuteRegionOp execute) -> LogicalResult { + return applyScfRegion(execute.getRegion(), execute.getResults(), + walk, state, execute); }) .template Case([&](scf::ForOp forOp) -> LogicalResult { - if constexpr (!std::is_same_v) { - return forOp.emitError() - << "scf.for is not supported for QCO DD functionality " - "construction"; - } else { - auto range = - resolveLoop(forOp, *walk.classical, walk.remainingExecutionSteps); - if (failed(range)) { - return failure(); - } + auto range = + resolveLoop(forOp, *walk.classical, walk.remainingExecutionSteps); + if (failed(range)) { + return failure(); + } - Block& body = *forOp.getBody(); - SmallVector carried(forOp.getInits().begin(), - forOp.getInits().end()); + Block& body = *forOp.getBody(); + SmallVector carried(forOp.getInits().begin(), + forOp.getInits().end()); - for (size_t t = 0; t < range->trips; - ++t, range->induction += range->step) { - if (walk.remainingExecutionSteps == 0) { - return forOp.emitError( - "QCO DD execution exceeds the limit of 10000 control-flow " - "steps"); - } - --walk.remainingExecutionSteps; - auto iterArgs = body.getArguments().drop_front(); - if (failed(bindValuePairs(carried, iterArgs, walk, forOp))) { - return failure(); - } - bindInteger( - body.getArgument(0), - range->induction.trunc(range->induction.getBitWidth() - 1), - *walk.classical); - if (failed(walkBlock(body, walk, state))) { - return failure(); - } - auto yield = cast(body.getTerminator()); - carried.assign(yield.getOperands().begin(), - yield.getOperands().end()); + for (size_t t = 0; t < range->trips; + ++t, range->induction += range->step) { + if (walk.remainingExecutionSteps == 0) { + return forOp.emitError( + "QCO DD execution exceeds the limit of 10000 control-flow " + "steps"); + } + --walk.remainingExecutionSteps; + auto iterArgs = body.getArguments().drop_front(); + if (failed(bindValuePairs(carried, iterArgs, walk, forOp))) { + return failure(); } - return bindValuePairs(carried, forOp.getResults(), walk, forOp); + if (failed(bindInteger( + body.getArgument(0), + range->induction.trunc(range->induction.getBitWidth() - 1), + *walk.classical))) { + return failure(); + } + if (failed(walkBlock(body, walk, state))) { + return failure(); + } + auto yield = cast(body.getTerminator()); + carried.assign(yield.getOperands().begin(), + yield.getOperands().end()); + } + return bindValuePairs(carried, forOp.getResults(), walk, forOp); + }) + .template Case([&](scf::WhileOp whileOp) -> LogicalResult { + if (!whileOp.getBefore().hasOneBlock() || + !whileOp.getAfter().hasOneBlock()) { + return whileOp.emitError() + << "scf.while regions must contain one block"; + } + Block& before = whileOp.getBefore().front(); + Block& after = whileOp.getAfter().front(); + SmallVector carried(whileOp.getInits().begin(), + whileOp.getInits().end()); + while (true) { + if (failed(bindValuePairs(carried, before.getArguments(), walk, + whileOp)) || + failed(walkBlock(before, walk, state))) { + return failure(); + } + auto condition = dyn_cast(before.getTerminator()); + if (!condition) { + return whileOp.emitError() + << "scf.while before region missing scf.condition"; + } + auto value = + lookupBool(condition.getCondition(), *walk.classical, whileOp); + if (failed(value)) { + return failure(); + } + if (!*value) { + return bindValuePairs(condition.getArgs(), whileOp.getResults(), + walk, whileOp); + } + if (walk.remainingExecutionSteps == 0) { + return whileOp.emitError( + "QCO DD execution exceeds the limit of 10000 control-flow " + "steps"); + } + --walk.remainingExecutionSteps; + if (failed(bindValuePairs(condition.getArgs(), after.getArguments(), + walk, whileOp)) || + failed(walkBlock(after, walk, state))) { + return failure(); + } + auto yield = dyn_cast(after.getTerminator()); + if (!yield) { + return whileOp.emitError() + << "scf.while after region missing scf.yield"; + } + carried.assign(yield.getOperands().begin(), + yield.getOperands().end()); } }) .template Case([&](func::CallOp call) -> LogicalResult { @@ -864,7 +1658,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { .template Case([&](CtrlOp ctrlOp) -> LogicalResult { if (auto inner = mqt::getSoleBodyUnitary( *ctrlOp.getBody())) { - auto decoded = decodeStandardGate(inner); + auto decoded = decodeStandardGate(inner, *walk.classical); if (failed(decoded)) { return failure(); } @@ -886,7 +1680,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { }) .template Case( [&](UnitaryOpInterface unitary) -> LogicalResult { - auto decoded = decodeStandardGate(unitary); + auto decoded = decodeStandardGate(unitary, *walk.classical); if (failed(decoded)) { return failure(); } @@ -909,6 +1703,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { template static LogicalResult walkFunction(func::FuncOp func, WalkState& walkState, StateDD& state) { + walkState.activeCalls.insert(func.getOperation()); // Function bodies include `func.return` as terminator; region walks skip // `qco.yield` and bind it separately. for (Operation& op : func.getBody().front()) { @@ -919,13 +1714,24 @@ static LogicalResult walkFunction(func::FuncOp func, WalkState& walkState, return success(); } -static FailureOr prepare(func::FuncOp func, const dd::Package& dd) { +namespace { +struct PreparedState { + QubitMap qubits; + TensorMap tensors; +}; +} // namespace + +static FailureOr prepare(func::FuncOp func, + const dd::Package& dd, + const DDBindings& bindings, + bool bindEntryAllocations = false) { if (!func.getBody().hasOneBlock()) { return func.emitError() << "QCO DD construction expects a single-block function body"; } - QubitMap qubits; + PreparedState prepared; + QubitMap& qubits = prepared.qubits; for (StaticOp staticOp : func.getBody().front().getOps()) { const auto q = static_cast(staticOp.getIndex()); qubits.bind(staticOp.getQubit(), q); @@ -934,32 +1740,63 @@ static FailureOr prepare(func::FuncOp func, const dd::Package& dd) { if (qubits.numQubits == 0) { qc::Qubit next = 0; for (Value arg : func.getArguments()) { - if (!isa(arg.getType())) { - continue; + if (isa(arg.getType())) { + qubits.bind(arg, next++); + } else if (isQTensorType(arg.getType())) { + const auto type = cast(arg.getType()); + int64_t size = type.getDimSize(0); + if (type.isDynamicDim(0)) { + const auto binding = bindings.find(arg); + if (binding == bindings.end() || !isa(binding->second)) { + return func.emitError() + << "dynamic qtensor arguments require an integer extent"; + } + size = cast(binding->second).getInt(); + if (size < 0) { + return func.emitError() + << "dynamic qtensor extent must be non-negative"; + } + } + TensorSlots slots; + slots.reserve(static_cast(size)); + for (int64_t i = 0; i < size; ++i) { + slots.emplace_back(next++); + } + prepared.tensors.bind(arg, std::move(slots)); } - qubits.bind(arg, next++); } - qubits.numQubits = next; + qubits.numQubits = static_cast(next); } - for (AllocOp alloc : func.getBody().front().getOps()) { - qubits.bind(alloc.getResult(), static_cast(qubits.numQubits++)); + if (bindEntryAllocations) { + for (AllocOp alloc : func.getBody().front().getOps()) { + qubits.bind(alloc.getResult(), + static_cast(qubits.numQubits++)); + } } if (dd.qubits() < qubits.numQubits) { return func.emitError() << "DD package has " << dd.qubits() << " qubits but function uses " << qubits.numQubits; } - return qubits; + return prepared; } -FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd) { - auto qubitsOr = prepare(func, dd); - if (failed(qubitsOr)) { +FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd, + const DDBindings& bindings) { + auto prepared = prepare(func, dd, bindings, /*bindEntryAllocations=*/true); + if (failed(prepared)) { return failure(); } - QubitMap qubits = std::move(*qubitsOr); + QubitMap qubits = std::move(prepared->qubits); + TensorMap tensors = std::move(prepared->tensors); ClassicalEnv classical; - WalkState walkState{ - .qubits = &qubits, .classical = &classical, .dd = &dd, .rng = nullptr}; + if (failed(applyBindings(func, bindings, classical))) { + return failure(); + } + WalkState walkState{.qubits = &qubits, + .tensors = &tensors, + .classical = &classical, + .dd = &dd, + .rng = nullptr}; dd::MatrixDD state = qubits.numQubits == 0 @@ -976,20 +1813,28 @@ FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd) { static FailureOr simulateImpl(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, - const QubitMap& preparedQubits, std::mt19937_64* rng, + const PreparedState& prepared, std::mt19937_64* rng, + const DDBindings& bindings, const DenseSet* deferredMeasurements = nullptr, ClassicalEnv* finalClassical = nullptr) { const size_t inputQubits = in.isTerminal() ? 0U : static_cast(in.p->v) + 1U; - if (inputQubits < preparedQubits.numQubits) { + if (inputQubits < prepared.qubits.numQubits) { dd.decRef(in); return func.emitError() << "input state has " << inputQubits << " qubits but function uses " - << preparedQubits.numQubits; + << prepared.qubits.numQubits; } - QubitMap qubits = preparedQubits; + QubitMap qubits = prepared.qubits; + qubits.numQubits = inputQubits; + TensorMap tensors = prepared.tensors; ClassicalEnv classical; + if (failed(applyBindings(func, bindings, classical))) { + dd.decRef(in); + return failure(); + } WalkState walkState{.qubits = &qubits, + .tensors = &tensors, .classical = &classical, .dd = &dd, .rng = rng, @@ -1007,13 +1852,14 @@ simulateImpl(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, } FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, - dd::Package& dd, std::mt19937_64& rng) { - auto qubits = prepare(func, dd); - if (failed(qubits)) { + dd::Package& dd, std::mt19937_64& rng, + const DDBindings& bindings) { + auto prepared = prepare(func, dd, bindings); + if (failed(prepared)) { dd.decRef(in); return failure(); } - return simulateImpl(func, in, dd, *qubits, &rng); + return simulateImpl(func, in, dd, *prepared, &rng, bindings); } static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { @@ -1094,7 +1940,7 @@ static FailureOr getSamplingPlan(func::FuncOp func) { static FailureOr encodeOutcome(ArrayRef outputs, const ClassicalEnv& classical, - StringRef basis, size_t numQubits) { + StringRef basis) { if (outputs.empty()) { return basis.str(); } @@ -1110,9 +1956,8 @@ static FailureOr encodeOutcome(ArrayRef outputs, const auto& cell = (*reg->second)[index]; if (cell.value) { outcome.push_back(*cell.value ? '1' : '0'); - } else if (cell.deferredWire && basis.size() == numQubits && - *cell.deferredWire < numQubits) { - outcome.push_back(basis[numQubits - 1 - *cell.deferredWire]); + } else if (cell.deferredWire && *cell.deferredWire < basis.size()) { + outcome.push_back(basis[basis.size() - 1 - *cell.deferredWire]); } else { return emitError(value.getLoc()) << "returned CBit register element " << index << " is undefined"; @@ -1122,10 +1967,12 @@ static FailureOr encodeOutcome(ArrayRef outputs, return outcome; } -FailureOr> -sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng) { - auto qubits = prepare(func, dd); - if (failed(qubits)) { +FailureOr> sample(func::FuncOp func, + dd::Package& dd, size_t shots, + std::mt19937_64& rng, + const DDBindings& bindings) { + auto prepared = prepare(func, dd, bindings); + if (failed(prepared)) { return failure(); } auto plan = getSamplingPlan(func); @@ -1138,10 +1985,10 @@ sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng) { return counts; } - const size_t numQubits = qubits->numQubits; + const size_t numQubits = prepared->qubits.numQubits; const auto record = [&](const ClassicalEnv& classical, StringRef basis) -> LogicalResult { - auto outcome = encodeOutcome(plan->outputs, classical, basis, numQubits); + auto outcome = encodeOutcome(plan->outputs, classical, basis); if (failed(outcome)) { return failure(); } @@ -1151,9 +1998,9 @@ sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng) { if (!plan->dynamic) { ClassicalEnv classical; - auto state = - simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, *qubits, - nullptr, &plan->deferredMeasurements, &classical); + auto state = simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, + *prepared, nullptr, bindings, + &plan->deferredMeasurements, &classical); if (failed(state)) { return failure(); } @@ -1169,7 +2016,7 @@ sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng) { for (size_t i = 0; i < shots; ++i) { ClassicalEnv classical; auto state = simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, - *qubits, &rng, nullptr, &classical); + *prepared, &rng, bindings, nullptr, &classical); if (failed(state)) { return failure(); } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 72bda4816f..e1c647047f 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -21,11 +21,13 @@ #include "mlir/Dialect/QCO/Builder/QCOProgramBuilder.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/Utils/DDFunctionality.h" +#include "mlir/Dialect/QTensor/IR/QTensorDialect.h" #include #include #include #include +#include #include #include #include @@ -62,8 +64,9 @@ class QCODDFunctionalityTest : public testing::Test { void SetUp() override { DialectRegistry registry; - registry.insert(); + registry.insert(); context = std::make_unique(); context->appendDialectRegistry(registry); context->loadAllAvailableDialects(); @@ -564,6 +567,40 @@ TEST_F(QCODDFunctionalityTest, SimulationConsumesInputReference) { EXPECT_TRUE(zeroQubitDd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, + SimulationPreservesWiderInputAcrossRuntimeAllocation) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main(%q0: !qco.qubit) { + %q1 = qco.alloc : !qco.qubit + %q2 = qco.x %q1 : !qco.qubit -> !qco.qubit + qco.sink %q0 : !qco.qubit + qco.sink %q2 : !qco.qubit + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(3); + auto input = dd->applyOperation( + dd->makeGateDD(dd::opToSingleQubitGateMatrix(qc::OpType::X), 1), + dd::makeZeroState(2, *dd)); + const auto output = simulate(mainFunc(*mod), input, *dd, rng); + ASSERT_TRUE(succeeded(output)); + + auto expected = dd->applyOperation( + dd->makeGateDD(dd::opToSingleQubitGateMatrix(qc::OpType::X), 1), + dd::makeZeroState(3, *dd)); + expected = dd->applyOperation( + dd->makeGateDD(dd::opToSingleQubitGateMatrix(qc::OpType::X), 2), + expected); + EXPECT_EQ(output->getVector(), expected.getVector()); + dd->decRef(*output); + dd->decRef(expected); +} + TEST_F(QCODDFunctionalityTest, SimulateMeasureCollapsesLikePackage) { auto mod = buildModule([](QCOProgramBuilder& b) { auto q = b.h(b.staticQubit(0)); @@ -632,29 +669,12 @@ TEST_F(QCODDFunctionalityTest, SimulateIfConstantBranches) { ASSERT_TRUE(thenMod); ASSERT_TRUE(elseMod); - auto dd = std::make_unique(1); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*thenMod), *dd))); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*elseMod), *dd))); - std::mt19937_64 rng(0); - auto zero = dd::makeZeroState(1, *dd); - auto one = dd->applyOperation( - dd->makeGateDD(dd::opToSingleQubitGateMatrix(qc::OpType::X), 0), - dd::makeZeroState(1, *dd)); - - const auto thenOut = - simulate(mainFunc(*thenMod), dd::makeZeroState(1, *dd), *dd, rng); - ASSERT_TRUE(succeeded(thenOut)); - EXPECT_EQ(thenOut->getVector(), one.getVector()); + qc::QuantumComputation thenQc(1); + thenQc.x(0); + expectEqualToQc(mainFunc(*thenMod), thenQc); - const auto elseOut = - simulate(mainFunc(*elseMod), dd::makeZeroState(1, *dd), *dd, rng); - ASSERT_TRUE(succeeded(elseOut)); - EXPECT_EQ(elseOut->getVector(), zero.getVector()); - - dd->decRef(*thenOut); - dd->decRef(*elseOut); - dd->decRef(zero); - dd->decRef(one); + const qc::QuantumComputation elseQc(1); + expectEqualToQc(mainFunc(*elseMod), elseQc); } TEST_F(QCODDFunctionalityTest, SimulateIndexSwitchBranches) { @@ -681,29 +701,12 @@ TEST_F(QCODDFunctionalityTest, SimulateIndexSwitchBranches) { ASSERT_TRUE(caseMod); ASSERT_TRUE(defaultMod); - auto dd = std::make_unique(1); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*caseMod), *dd))); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*defaultMod), *dd))); - std::mt19937_64 rng(0); - auto zero = dd::makeZeroState(1, *dd); - auto one = dd->applyOperation( - dd->makeGateDD(dd::opToSingleQubitGateMatrix(qc::OpType::X), 0), - dd::makeZeroState(1, *dd)); - - const auto caseOut = - simulate(mainFunc(*caseMod), dd::makeZeroState(1, *dd), *dd, rng); - ASSERT_TRUE(succeeded(caseOut)); - EXPECT_EQ(caseOut->getVector(), one.getVector()); + qc::QuantumComputation caseQc(1); + caseQc.x(0); + expectEqualToQc(mainFunc(*caseMod), caseQc); - const auto defaultOut = - simulate(mainFunc(*defaultMod), dd::makeZeroState(1, *dd), *dd, rng); - ASSERT_TRUE(succeeded(defaultOut)); - EXPECT_EQ(defaultOut->getVector(), zero.getVector()); - - dd->decRef(*caseOut); - dd->decRef(*defaultOut); - dd->decRef(zero); - dd->decRef(one); + const qc::QuantumComputation defaultQc(1); + expectEqualToQc(mainFunc(*defaultMod), defaultQc); } TEST_F(QCODDFunctionalityTest, SimulateMeasureFeedsIf) { @@ -1137,6 +1140,64 @@ TEST_F(QCODDFunctionalityTest, EmbedsWideLocalMatrixWithoutRegisterLimit) { expectEqualToQc(mainFunc(*mod), qc); } +TEST_F(QCODDFunctionalityTest, RejectsUnsupportedOrUnboundClassicalOperations) { + for (const StringRef source : { + R"mlir(module { + func.func @main(%c: i1) { + %q = qco.static 0 : !qco.qubit + %bad = arith.index_castui %c : i1 to index + qco.sink %q : !qco.qubit + return + } + })mlir", + R"mlir(module { + func.func @main(%unmapped: i1) { + %q = qco.static 0 : !qco.qubit + %true = arith.constant true + %bad = arith.andi %unmapped, %true : i1 + qco.sink %q : !qco.qubit + return + } + })mlir", + R"mlir(module { + func.func @main(%unmapped: index) { + %q = qco.static 0 : !qco.qubit + %one = arith.constant 1 : index + %bad = arith.ori %unmapped, %one : index + qco.sink %q : !qco.qubit + return + } + })mlir", + R"mlir(module { + func.func @main() { + %bad = arith.constant 1.0 : f32 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %one = arith.constant 1 : i32 + %bad = arith.sitofp %one : i32 to f32 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %q = qco.static 0 : !qco.qubit + %one = arith.constant 1 : i64 + %bad = arith.maxsi %one, %one : i64 + qco.sink %q : !qco.qubit + return + } + })mlir"}) { + auto mod = parseSourceString(source, context.get()); + 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, RejectsUnmappedClassicalControl) { for (const StringRef source : {R"mlir( module { @@ -1254,6 +1315,45 @@ TEST_F(QCODDFunctionalityTest, BindsClassicalIndexResults) { dd->decRef(expected); } +TEST_F(QCODDFunctionalityTest, RejectsUnboundClassicalRegionResults) { + for (const StringRef source : { + R"mlir(module { + func.func @main(%unmapped: i1) { + %q = qco.static 0 : !qco.qubit + %true = arith.constant true + %result, %out = qco.if %true args(%arg = %q) + -> (i1, !qco.qubit) { + qco.yield %unmapped, %arg : i1, !qco.qubit + } else args(%arg = %q) { + qco.yield %true, %arg : i1, !qco.qubit + } + qco.sink %out : !qco.qubit + return + } + })mlir", + R"mlir(module { + func.func @main(%unmapped: index) { + %q = qco.static 0 : !qco.qubit + %true = arith.constant true + %zero = arith.constant 0 : index + %result, %out = qco.if %true args(%arg = %q) + -> (index, !qco.qubit) { + qco.yield %unmapped, %arg : index, !qco.qubit + } else args(%arg = %q) { + qco.yield %zero, %arg : index, !qco.qubit + } + qco.sink %out : !qco.qubit + return + } + })mlir"}) { + auto mod = parseSourceString(source, context.get()); + 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, Rejects) { { auto mod = buildModule([](QCOProgramBuilder& b) { @@ -1596,6 +1696,27 @@ TEST_F(QCODDFunctionalityTest, ScfForSnapshotsYieldedInductionValue) { } TEST_F(QCODDFunctionalityTest, RejectsUnsupportedFuncCalls) { + auto selfRecursive = parseSourceString(R"mlir( + module { + func.func @main(%recurse: i1) { + scf.if %recurse { + %false = arith.constant false + func.call @main(%false) : (i1) -> () + } + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(selfRecursive); + auto selfRecursiveFunc = mainFunc(*selfRecursive); + DDBindings bindings; + bindings[selfRecursiveFunc.getArgument(0)] = + BoolAttr::get(context.get(), true); + auto zeroQubitDd = std::make_unique(0); + EXPECT_TRUE(failed(simulate(selfRecursiveFunc, dd::VectorDD::one(), + *zeroQubitDd, rng, bindings))); + auto recursive = parseSourceString(R"mlir( module { func.func @rec(%q: !qco.qubit) -> !qco.qubit { @@ -1853,6 +1974,24 @@ TEST_F(QCODDFunctionalityTest, EXPECT_TRUE(dd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, SampleDefersAllocatedQubitMeasurement) { + auto mod = buildModule([](QCOProgramBuilder& b) { + auto reg = + b.allocClassicalBitRegister(1, {}, cbit::Initialization::Undefined); + auto q = b.x(b.allocQubit()); + std::tie(q, std::ignore) = b.measure(q, reg, 0); + b.sink(q); + return reg; + }); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(1); + std::mt19937_64 rng(11); + const auto histogram = sample(mainFunc(*mod), *dd, 8, rng); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"1", 8}})); +} + TEST_F(QCODDFunctionalityTest, SampleExecutesControlMeasurementPerShot) { auto mod = buildModule([](QCOProgramBuilder& b) { auto reg = @@ -1915,6 +2054,47 @@ TEST_F(QCODDFunctionalityTest, FuncCallSharesClassicalCBitStorage) { expectSimulatesFromZero(mainFunc(*mod), true); } +TEST_F(QCODDFunctionalityTest, RejectsUnsupportedClassicalMemRefs) { + for (const StringRef source : { + R"mlir(module { + func.func @main(%reg: memref) { + %value = memref.load %reg[] : memref + return + } + })mlir", + R"mlir(module { + func.func @main(%reg: memref) { + %value = arith.constant true + memref.store %value, %reg[] : memref + return + } + })mlir", + R"mlir(module { + func.func @main(%n: index) { + %reg = memref.alloc(%n) : memref + memref.dealloc %reg : memref + return + } + })mlir", + R"mlir(module { + func.func @main() { + %reg = memref.alloc() : memref<1xf32> + memref.dealloc %reg : memref<1xf32> + return + } + })mlir", + R"mlir(module { + func.func @main() { + %reg = memref.alloc() : memref<1xi1> + %value = arith.constant true + %i2 = arith.constant 2 : index + memref.store %value, %reg[%i2] : memref<1xi1> + return + } + })mlir"}) { + expectMlirSimulationFails(0, source); + } +} TEST_F(QCODDFunctionalityTest, SampleExecutesCalleeMeasurementBeforeCallerGate) { auto mod = parseSourceString(R"mlir( @@ -1993,4 +2173,324 @@ TEST_F(QCODDFunctionalityTest, SampleExecutesNestedMeasurementPerShot) { EXPECT_TRUE(dd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, SymbolicParametersUseBindings) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main(%theta: f64) { + %q = qco.static 0 : !qco.qubit + %twice = arith.addf %theta, %theta : f64 + %q1 = qco.rx(%twice) %q : !qco.qubit -> !qco.qubit + qco.gphase(%theta) + qco.sink %q1 : !qco.qubit + return + } + } + )mlir", + context.get()); + auto concrete = buildModule([](QCOProgramBuilder& b) { + auto q = b.rx(std::numbers::pi, b.staticQubit(0)); + b.gphase(std::numbers::pi / 2.0); + b.sink(q); + return b.intConstant(0); + }); + ASSERT_TRUE(mod); + ASSERT_TRUE(concrete); + + auto func = mainFunc(*mod); + DDBindings bindings; + bindings[func.getArgument(0)] = FloatAttr::get( + cast(func.getArgument(0).getType()), std::numbers::pi / 2.0); + + auto dd = std::make_unique(1); + auto actual = buildFunctionality(func, *dd, bindings); + auto expected = buildFunctionality(mainFunc(*concrete), *dd); + ASSERT_TRUE(succeeded(actual)); + ASSERT_TRUE(succeeded(expected)); + EXPECT_EQ(actual->getMatrix(1), expected->getMatrix(1)); + dd->decRef(*actual); + dd->decRef(*expected); + + std::mt19937_64 rng(5); + const auto histogram = sample(func, *dd, 8, rng, bindings); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"1", 8}})); + + EXPECT_TRUE(failed(buildFunctionality(func, *dd))); + bindings[func.getArgument(0)] = + IntegerAttr::get(IntegerType::get(context.get(), 64), 1); + EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); + bindings[func.getArgument(0)] = + FloatAttr::get(Float32Type::get(context.get()), 1.0); + EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); +} + +TEST_F(QCODDFunctionalityTest, BuildsThroughConcreteControlFlow) { + auto mod = buildModule([](QCOProgramBuilder& b) { + auto q = b.staticQubit(0); + q = b.qcoIf( + true, q, [&](Value arg) { return b.x(arg); }, + [&](Value arg) { return arg; }); + q = b.qcoIndexSwitch(1, q, ArrayRef{0, 1}, + SmallVector>{ + [&](Value arg) { return b.h(arg); }, + [&](Value arg) { return b.z(arg); }}, + [&](Value arg) { return arg; }); + q = b.scfFor(0, 2, 1, ValueRange{q.value}, + [&](Value /*index*/, ValueRange args) -> SmallVector { + return {b.h(args[0])}; + })[0]; + b.sink(q); + return b.intConstant(0); + }); + ASSERT_TRUE(mod); + + qc::QuantumComputation qc(1); + qc.x(0); + qc.z(0); + qc.h(0); + qc.h(0); + expectEqualToQc(mainFunc(*mod), qc); +} + +TEST_F(QCODDFunctionalityTest, StandardScfRegionsAndWhileCarryValues) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main() { + %q = qco.static 0 : !qco.qubit + %q1 = scf.execute_region -> !qco.qubit { + %out = qco.x %q : !qco.qubit -> !qco.qubit + scf.yield %out : !qco.qubit + } + %true = arith.constant true + %selector = scf.if %true -> index { + %one = arith.constant 1 : index + scf.yield %one : index + } else { + %zero = arith.constant 0 : index + scf.yield %zero : index + } + %apply_z = scf.index_switch %selector -> i1 + case 1 { + %yes = arith.constant true + scf.yield %yes : i1 + } + default { + %no = arith.constant false + scf.yield %no : i1 + } + %q2 = qco.if %apply_z args(%qarg = %q1) -> (!qco.qubit) { + %out = qco.z %qarg : !qco.qubit -> !qco.qubit + qco.yield %out : !qco.qubit + } else args(%qarg = %q1) { + qco.yield %qarg : !qco.qubit + } + %zero = arith.constant 0 : index + %result:2 = scf.while (%qarg = %q2, %i = %zero) + : (!qco.qubit, index) -> (!qco.qubit, index) { + %one = arith.constant 1 : index + %condition = arith.cmpi slt, %i, %one : index + scf.condition(%condition) %qarg, %i : !qco.qubit, index + } do { + ^bb0(%qarg: !qco.qubit, %i: index): + %out = qco.x %qarg : !qco.qubit -> !qco.qubit + %one = arith.constant 1 : index + %next = arith.addi %i, %one : index + scf.yield %out, %next : !qco.qubit, index + } + %false = arith.constant false + %final = scf.while (%qarg = %result#0) + : (!qco.qubit) -> !qco.qubit { + scf.condition(%false) %qarg : !qco.qubit + } do { + ^bb0(%qarg: !qco.qubit): + %unreachable = qco.h %qarg : !qco.qubit -> !qco.qubit + scf.yield %unreachable : !qco.qubit + } + qco.sink %final : !qco.qubit + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + qc::QuantumComputation qc(1); + qc.x(0); + qc.z(0); + qc.x(0); + expectEqualToQc(mainFunc(*mod), qc); +} + +TEST_F(QCODDFunctionalityTest, DynamicAllocationsAndQTensorBookkeeping) { + auto mod = buildModule([](QCOProgramBuilder& b) { + auto q0 = b.x(b.allocQubit()); + auto one = arith::ConstantIndexOp::create(b, 1).getResult(); + auto tensor = b.qtensorAlloc(one); + Value remaining; + Value q1; + std::tie(remaining, q1) = b.qtensorExtract(tensor, 0); + auto output = b.qtensorFromElements({q0, b.x(q1)}); + b.qtensorDealloc(remaining); + b.qtensorDealloc(output); + return b.intConstant(0); + }); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(2); + std::mt19937_64 rng(3); + const auto histogram = sample(mainFunc(*mod), *dd, 8, rng); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"11", 8}})); + + auto smallDd = std::make_unique(1); + EXPECT_TRUE(failed(sample(mainFunc(*mod), *smallDd, 1, rng))); + + auto invalidIndex = parseSourceString(R"mlir( + module { + func.func @main() { + %one = arith.constant 1 : index + %tensor = qtensor.alloc(%one) : tensor + %remaining, %q = qtensor.extract %tensor[%one] + : tensor + qco.sink %q : !qco.qubit + qtensor.dealloc %remaining : tensor + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(invalidIndex); + auto oneQubitDd = std::make_unique(1); + EXPECT_TRUE(failed(sample(mainFunc(*invalidIndex), *oneQubitDd, 1, rng))); +} + +TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main(%arg0: tensor) + -> tensor { + %one = arith.constant 1 : index + %remaining, %q = qtensor.extract %arg0[%one] + : tensor + %q1 = qco.x %q : !qco.qubit -> !qco.qubit + %result = qtensor.insert %q1 into %remaining[%one] + : tensor + return %result : tensor + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + auto func = mainFunc(*mod); + DDBindings bindings; + bindings[func.getArgument(0)] = + IntegerAttr::get(IndexType::get(context.get()), 2); + + auto dd = std::make_unique(2); + std::mt19937_64 rng(7); + const auto histogram = sample(func, *dd, 4, rng, bindings); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"10", 4}})); + + EXPECT_TRUE(failed(buildFunctionality(func, *dd))); + bindings[func.getArgument(0)] = + IntegerAttr::get(IndexType::get(context.get()), -1); + EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); +} + +TEST_F(QCODDFunctionalityTest, QTensorFlowsThroughLoopAndCall) { + auto mod = parseSourceString(R"mlir( + module { + func.func @flip(%arg: tensor<1x!qco.qubit>) + -> tensor<1x!qco.qubit> { + %zero = arith.constant 0 : index + %remaining, %q = qtensor.extract %arg[%zero] + : tensor<1x!qco.qubit> + %q1 = qco.x %q : !qco.qubit -> !qco.qubit + %result = qtensor.insert %q1 into %remaining[%zero] + : tensor<1x!qco.qubit> + return %result : tensor<1x!qco.qubit> + } + func.func @main() { + %zero = arith.constant 0 : index + %one = arith.constant 1 : index + %tensor = qtensor.alloc(%one) : tensor<1x!qco.qubit> + %result = scf.for %i = %zero to %one step %one + iter_args(%arg = %tensor) -> tensor<1x!qco.qubit> { + %next = func.call @flip(%arg) + : (tensor<1x!qco.qubit>) -> tensor<1x!qco.qubit> + scf.yield %next : tensor<1x!qco.qubit> + } + qtensor.dealloc %result : tensor<1x!qco.qubit> + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(1); + std::mt19937_64 rng(13); + const auto histogram = sample(mainFunc(*mod), *dd, 4, rng); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"1", 4}})); +} + +TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { + auto mod = parseSourceString(R"mlir( + module { + func.func @set(%reg: memref, %value: i16) { + %zero = arith.constant 0 : index + memref.store %value, %reg[%zero] : memref + return + } + func.func @main() { + %one = arith.constant 1 : index + %reg = memref.alloc(%one) : memref + %three = arith.constant 3 : i16 + %four = arith.constant 4 : i16 + %seven = arith.addi %three, %four : i16 + %two = arith.constant 2 : i16 + %fourteen = arith.muli %seven, %two : i16 + %quotient = arith.divsi %fourteen, %two : i16 + %remainder = arith.remui %quotient, %two : i16 + %shifted = arith.shli %remainder, %two : i16 + %restored = arith.shrui %shifted, %two : i16 + %wide = arith.extui %restored : i16 to i32 + %narrow = arith.trunci %wide : i32 to i16 + %as_float = arith.sitofp %narrow : i16 to f64 + %back = arith.fptosi %as_float : f64 to i16 + func.call @set(%reg, %quotient) : (memref, i16) -> () + %zero = arith.constant 0 : index + %stored = memref.load %reg[%zero] : memref + %expected = arith.constant 7 : i16 + %integer_ok = arith.cmpi eq, %stored, %expected : i16 + %casts_ok = arith.cmpi eq, %back, %remainder : i16 + %one_float = arith.constant 1.0 : f64 + %two_float = arith.addf %one_float, %one_float : f64 + %four_float = arith.addf %two_float, %two_float : f64 + %half = arith.divf %four_float, %two_float : f64 + %float_remainder = arith.remf %half, %one_float : f64 + %zero_float = arith.constant 0.0 : f64 + %float_ok = arith.cmpf oeq, %float_remainder, %zero_float : f64 + %integer_and_casts = arith.andi %integer_ok, %casts_ok : i1 + %condition = arith.andi %integer_and_casts, %float_ok : i1 + %q = qco.static 0 : !qco.qubit + %q1 = qco.if %condition args(%qin = %q) -> (!qco.qubit) { + %out = qco.x %qin : !qco.qubit -> !qco.qubit + qco.yield %out : !qco.qubit + } else args(%qin = %q) { + qco.yield %qin : !qco.qubit + } + memref.dealloc %reg : memref + qco.sink %q1 : !qco.qubit + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + expectSimulatesFromZero(mainFunc(*mod), true); +} + } // namespace From 93c7050b0f40accfc1c0a73ffc754e117d48c33f Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 16:05:56 +0200 Subject: [PATCH 02/12] =?UTF-8?q?=F0=9F=93=9D=20Clarify=20QCO=20DD=20inter?= =?UTF-8?q?preter=20documentation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- CHANGELOG.md | 2 +- mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h | 9 ++++++--- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 252867bf49..55b7ec71d7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -38,7 +38,7 @@ releases may include breaking changes. `f64` SSA, CBit registers, one-dimensional memrefs, dynamic quantum allocation and qtensors, dense multi-wire embedding, output-aware multi-shot sampling, and Python bindings ([#1915], [#1973], [#2077], [#2078]) - ([**@simon1hofmann**]) + ([**@simon1hofmann**], [**@burgholzer**]) - ✨ Add immutable MLIR compiler targets, QDMI device integration, and target compilation through C++, Python, and `mqt-cc` ([#1687], [#1993], [#1999], [#2049]) ([**@MatthiasReumann**], [**@simon1hofmann**], [**@burgholzer**]) diff --git a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h index bb4fdf22cf..46925fb7e6 100644 --- a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h +++ b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h @@ -67,7 +67,8 @@ using DDBindings = DenseMap; * * @param func The QCO function to construct the functionality for * @param dd The DD package to use (must hold at least the function's qubits) - * @param bindings Concrete values for symbolic function arguments + * @param bindings Concrete scalar values and dynamic QTensor extents for entry + * arguments * @return The matrix DD on success, or failure for unsupported programs */ FailureOr @@ -99,7 +100,8 @@ buildFunctionality(func::FuncOp func, dd::Package& dd, * higher wires are preserved; one reference is consumed * @param dd The DD package to use * @param rng RNG used for collapsing measurements and resets - * @param bindings Concrete values for symbolic function arguments + * @param bindings Concrete scalar values and dynamic QTensor extents for entry + * arguments * @return The output statevector DD on success, or failure for unsupported * programs */ @@ -127,7 +129,8 @@ FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, * @param dd The DD package to use * @param shots Number of shots * @param rng RNG for collapsing measurements and non-collapsing sampling - * @param bindings Concrete values for symbolic function arguments + * @param bindings Concrete scalar values and dynamic QTensor extents for entry + * arguments * @return Histogram of outcome strings on success, or failure for unsupported * programs */ From b4743fa870f7204c4cc78c3fb6c222e951c923c6 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 16:06:17 +0200 Subject: [PATCH 03/12] =?UTF-8?q?=E2=9C=85=20Cover=20QCO=20DD=20interprete?= =?UTF-8?q?r=20branches?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- .../QCO/Utils/test_dd_functionality.cpp | 256 ++++++++++++++++++ 1 file changed, 256 insertions(+) diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index e1c647047f..69f6beb44a 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -2493,4 +2493,260 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { expectSimulatesFromZero(mainFunc(*mod), true); } +TEST_F(QCODDFunctionalityTest, + SupportsAdditionalClassicalOperationsAndBindings) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main(%idx: index, %word: i16, %flag: i1) { + %zero = arith.constant 0 : i8 + %one = arith.constant 1 : i8 + %two = arith.constant 2 : i8 + %four = arith.constant 4 : i8 + %negative = arith.constant -5 : i8 + %sle = arith.cmpi sle, %one, %two : i8 + %sgt = arith.cmpi sgt, %two, %one : i8 + %sge = arith.cmpi sge, %two, %two : i8 + %ult = arith.cmpi ult, %one, %two : i8 + %ule = arith.cmpi ule, %two, %two : i8 + %ugt = arith.cmpi ugt, %two, %one : i8 + %uge = arith.cmpi uge, %two, %two : i8 + %quotient = arith.divui %four, %two : i8 + %remainder = arith.remsi %negative, %two : i8 + %shifted = arith.shrsi %negative, %one : i8 + %extended = arith.extsi %negative : i8 to i16 + %as_index = arith.index_cast %word : i16 to index + %selected = arith.select %flag, %one, %zero : i8 + %one_float = arith.constant 1.0 : f64 + %two_float = arith.constant 2.0 : f64 + %difference = arith.subf %two_float, %one_float : f64 + %product = arith.mulf %difference, %two_float : f64 + %negated = arith.negf %product : f64 + %zero_index = arith.constant 0 : index + %indices = memref.alloc() : memref<1xindex> + %loaded_index = memref.load %indices[%zero_index] : memref<1xindex> + memref.dealloc %indices : memref<1xindex> + %floats = memref.alloc() : memref<1xf64> + %loaded_float = memref.load %floats[%zero_index] : memref<1xf64> + memref.dealloc %floats : memref<1xf64> + %false = arith.constant false + scf.if %false { + } + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + auto func = mainFunc(*mod); + DDBindings bindings; + bindings[func.getArgument(0)] = + IntegerAttr::get(IndexType::get(context.get()), 3); + bindings[func.getArgument(1)] = + IntegerAttr::get(IntegerType::get(context.get(), 16), -2); + bindings[func.getArgument(2)] = + IntegerAttr::get(IntegerType::get(context.get(), 1), 1); + auto dd = std::make_unique(0); + const auto output = simulate(func, dd::VectorDD::one(), *dd, rng, bindings); + ASSERT_TRUE(succeeded(output)); + EXPECT_TRUE(output->isTerminal()); + dd->decRef(*output); +} + +TEST_F(QCODDFunctionalityTest, RejectsClassicalRuntimeErrors) { + for (const StringRef source : { + R"mlir(module { + func.func @main(%unbound: i16) { + %true = arith.constant true + %zero = arith.constant 0 : i16 + %selected = arith.select %true, %unbound, %zero : i16 + return + } + })mlir", + R"mlir(module { + func.func @main(%reg: memref<1xi16>) { + %zero = arith.constant 0 : index + %value = memref.load %reg[%zero] : memref<1xi16> + return + } + })mlir", + R"mlir(module { + func.func @main(%index: index) { + %reg = memref.alloc() : memref<1xi16> + %value = memref.load %reg[%index] : memref<1xi16> + return + } + })mlir", + R"mlir(module { + func.func @main(%value: i16) { + %zero = arith.constant 0 : index + %reg = memref.alloc() : memref<1xi16> + memref.store %value, %reg[%zero] : memref<1xi16> + return + } + })mlir", + R"mlir(module { + func.func @main() { + %negative = arith.constant -1 : index + %reg = memref.alloc(%negative) : memref + return + } + })mlir", + R"mlir(module { + func.func @main() { + %zero = arith.constant 0 : i8 + %one = arith.constant 1 : i8 + %invalid = arith.divui %one, %zero : i8 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %huge = arith.constant 1.0e+300 : f64 + %invalid = arith.fptosi %huge : f64 to i8 + return + } + })mlir", + R"mlir(module { + func.func @main(%rhs: i8) { + %one = arith.constant 1 : i8 + %invalid = arith.divui %one, %rhs : i8 + return + } + })mlir", + R"mlir(module { + func.func @main(%lhs: i8) { + %one = arith.constant 1 : i8 + %invalid = arith.divui %lhs, %one : i8 + return + } + })mlir", + R"mlir(module { + func.func @main(%amount: i8) { + %one = arith.constant 1 : i8 + %invalid = arith.shli %one, %amount : i8 + return + } + })mlir", + R"mlir(module { + func.func @main(%lhs: f64) { + %zero = arith.constant 0.0 : f64 + %invalid = arith.cmpf oeq, %lhs, %zero : f64 + return + } + })mlir", + R"mlir(module { + func.func @main(%value: i8) { + %invalid = arith.sitofp %value : i8 to f64 + return + } + })mlir", + R"mlir(module { + func.func @main(%value: f64) { + %invalid = arith.fptosi %value : f64 to i8 + return + } + })mlir", + R"mlir(module { + func.func @consume(%reg: memref<1xi16>) { + return + } + func.func @main(%reg: memref<1xi16>) { + func.call @consume(%reg) : (memref<1xi16>) -> () + return + } + })mlir", + R"mlir(module { + func.func @main(%condition: i1) { + scf.if %condition { + } + return + } + })mlir", + R"mlir(module { + func.func @main(%selector: index) { + scf.index_switch %selector + default { + } + return + } + })mlir", + R"mlir(module { + func.func @main(%size: index) { + %tensor = qtensor.alloc(%size) : tensor + qtensor.dealloc %tensor : tensor + return + } + })mlir"}) { + expectMlirSimulationFails(0, source); + } + + auto mod = parseSourceString(R"mlir( + module { + func.func @main() { + %zero = arith.constant 0 : index + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + auto func = mainFunc(*mod); + auto constant = *func.getBody().front().getOps().begin(); + DDBindings bindings; + bindings[constant.getResult()] = + IntegerAttr::get(IndexType::get(context.get()), 0); + auto dd = std::make_unique(0); + EXPECT_TRUE(failed(simulate(func, dd::VectorDD::one(), *dd, rng, bindings))); + EXPECT_TRUE(dd->getRootSet().empty()); +} + +TEST_F(QCODDFunctionalityTest, BuildFunctionalityRestrictsRuntimeAllocations) { + auto topLevel = parseSourceString(R"mlir( + module { + func.func @main() { + %q = qco.alloc : !qco.qubit + %out = qco.x %q : !qco.qubit -> !qco.qubit + qco.sink %out : !qco.qubit + return + } + } + )mlir", + context.get()); + auto nested = parseSourceString(R"mlir( + module { + func.func @main() { + %true = arith.constant true + scf.if %true { + %q = qco.alloc : !qco.qubit + qco.sink %q : !qco.qubit + } + return + } + } + )mlir", + context.get()); + auto tensor = parseSourceString(R"mlir( + module { + func.func @main() { + %one = arith.constant 1 : index + %tensor = qtensor.alloc(%one) : tensor + qtensor.dealloc %tensor : tensor + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(topLevel); + ASSERT_TRUE(nested); + ASSERT_TRUE(tensor); + + auto dd = std::make_unique(1); + const auto functionality = buildFunctionality(mainFunc(*topLevel), *dd); + ASSERT_TRUE(succeeded(functionality)); + dd->decRef(*functionality); + EXPECT_TRUE(failed(buildFunctionality(mainFunc(*nested), *dd))); + EXPECT_TRUE(failed(buildFunctionality(mainFunc(*tensor), *dd))); +} + } // namespace From 1f5ad07dda081569e18c3679e3eec6675abe09d0 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 17:50:19 +0200 Subject: [PATCH 04/12] =?UTF-8?q?=F0=9F=90=9B=20Harden=20QCO=20DD=20prepar?= =?UTF-8?q?ation=20and=20sampling?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 124 ++++++++++++++---- .../QCO/Utils/test_dd_functionality.cpp | 48 +++++++ 2 files changed, 150 insertions(+), 22 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 318f11f73e..921edac1bd 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -1738,10 +1738,14 @@ static FailureOr prepare(func::FuncOp func, qubits.numQubits = std::max(qubits.numQubits, static_cast(q) + 1); } if (qubits.numQubits == 0) { - qc::Qubit next = 0; + size_t next = 0; for (Value arg : func.getArguments()) { if (isa(arg.getType())) { - qubits.bind(arg, next++); + if (next >= dd::Package::MAX_POSSIBLE_QUBITS) { + return func.emitError() + << "QCO function exceeds the supported qubit range"; + } + qubits.bind(arg, static_cast(next++)); } else if (isQTensorType(arg.getType())) { const auto type = cast(arg.getType()); int64_t size = type.getDimSize(0); @@ -1757,15 +1761,20 @@ static FailureOr prepare(func::FuncOp func, << "dynamic qtensor extent must be non-negative"; } } + const auto count = static_cast(size); + if (count > dd::Package::MAX_POSSIBLE_QUBITS - next) { + return func.emitError() + << "QCO function exceeds the supported qubit range"; + } TensorSlots slots; - slots.reserve(static_cast(size)); - for (int64_t i = 0; i < size; ++i) { - slots.emplace_back(next++); + slots.reserve(count); + for (size_t i = 0; i < count; ++i) { + slots.emplace_back(static_cast(next++)); } prepared.tensors.bind(arg, std::move(slots)); } } - qubits.numQubits = static_cast(next); + qubits.numQubits = next; } if (bindEntryAllocations) { for (AllocOp alloc : func.getBody().front().getOps()) { @@ -1872,21 +1881,91 @@ static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { }); } -static bool isDeferrableMeasurement(MeasureOp measure, Block* entry, +static bool hasOutputOnlyMeasurementResult(MeasureOp measure, + ArrayRef outputs) { + if (measure.getResult().use_empty()) { + return false; + } + return llvm::all_of(measure.getResult().getUses(), [&](const OpOperand& use) { + auto store = dyn_cast(use.getOwner()); + return store && isOutputOnlyRegister(store.getReg(), outputs); + }); +} + +static std::optional getConstantTensorIndex(Value value) { + const auto attr = mqt::valueToConstantAttr(value); + const auto integer = attr ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer || integer.getInt() < 0) { + return std::nullopt; + } + return static_cast(integer.getInt()); +} + +static bool hasOnlyTerminalQuantumUses(Value value, + std::optional tensorSlot, + func::FuncOp func, + ArrayRef outputs, + DenseSet& visited) { + if (!visited.insert(value).second) { + return true; + } + return llvm::all_of(value.getUses(), [&](OpOperand& use) { + Operation* owner = use.getOwner(); + if (isa(owner)) { + return true; + } + if (auto measure = dyn_cast(owner)) { + return !tensorSlot && measure->getParentOfType() == func && + hasOutputOnlyMeasurementResult(measure, outputs) && + hasOnlyTerminalQuantumUses(measure.getQubitOut(), std::nullopt, + func, outputs, visited); + } + if (auto fromElements = dyn_cast(owner)) { + return !tensorSlot && hasOnlyTerminalQuantumUses(fromElements.getResult(), + use.getOperandNumber(), + func, outputs, visited); + } + if (auto insert = dyn_cast(owner)) { + const auto index = getConstantTensorIndex(insert.getIndex()); + if (!index) { + return false; + } + if (use.get() == insert.getScalar()) { + return !tensorSlot && + hasOnlyTerminalQuantumUses(insert.getResult(), *index, func, + outputs, visited); + } + return tensorSlot && *tensorSlot != *index && + hasOnlyTerminalQuantumUses(insert.getResult(), tensorSlot, func, + outputs, visited); + } + if (auto extract = dyn_cast(owner)) { + const auto index = getConstantTensorIndex(extract.getIndex()); + if (!tensorSlot || !index) { + return false; + } + if (*tensorSlot == *index) { + return hasOnlyTerminalQuantumUses(extract.getResult(), std::nullopt, + func, outputs, visited); + } + return hasOnlyTerminalQuantumUses(extract.getOutTensor(), tensorSlot, + func, outputs, visited); + } + return false; + }); +} + +static bool isDeferrableMeasurement(MeasureOp measure, func::FuncOp func, ArrayRef outputs) { - return measure->getBlock() == entry && - llvm::all_of(measure.getQubitOut().getUses(), - [](const OpOperand& use) { - return isa(use.getOwner()); - }) && - llvm::all_of(measure.getResult().getUses(), [&](const OpOperand& use) { - auto store = dyn_cast(use.getOwner()); - return store && isOutputOnlyRegister(store.getReg(), outputs); - }); + if (!hasOutputOnlyMeasurementResult(measure, outputs)) { + return false; + } + DenseSet visited; + return hasOnlyTerminalQuantumUses(measure.getQubitOut(), std::nullopt, func, + outputs, visited); } -static void analyzeSampling(func::FuncOp func, Block* entry, - ArrayRef outputs, +static void analyzeSampling(func::FuncOp func, ArrayRef outputs, DenseSet& active, SamplingPlan& plan) { Operation* funcOp = func.getOperation(); if (!active.insert(funcOp).second) { @@ -1897,7 +1976,7 @@ static void analyzeSampling(func::FuncOp func, Block* entry, if (isa(op)) { plan.dynamic = true; } else if (auto measure = dyn_cast(op)) { - if (isDeferrableMeasurement(measure, entry, outputs)) { + if (isDeferrableMeasurement(measure, func, outputs)) { plan.deferredMeasurements.insert(op); } else { plan.dynamic = true; @@ -1905,10 +1984,11 @@ static void analyzeSampling(func::FuncOp func, Block* entry, } else if (auto call = dyn_cast(op)) { auto callee = SymbolTable::lookupNearestSymbolFrom( call, call.getCalleeAttr()); - if (!callee.getBody().hasOneBlock()) { + if (!callee || callee.isDeclaration() || + !callee.getBody().hasOneBlock()) { plan.dynamic = true; } else { - analyzeSampling(callee, entry, outputs, active, plan); + analyzeSampling(callee, outputs, active, plan); } } }); @@ -1934,7 +2014,7 @@ static FailureOr getSamplingPlan(func::FuncOp func) { } DenseSet active; - analyzeSampling(func, &entry, plan.outputs, active, plan); + analyzeSampling(func, plan.outputs, active, plan); return plan; } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 69f6beb44a..a56d6c9de4 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -1974,6 +1974,39 @@ TEST_F(QCODDFunctionalityTest, EXPECT_TRUE(dd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, DefersTensorMeasurementDespiteLaterUnrelatedOp) { + auto mod = buildModule([](QCOProgramBuilder& b) { + auto reg = + b.allocClassicalBitRegister(1, {}, cbit::Initialization::Undefined); + auto q0 = b.h(b.allocQubit()); + auto q1 = b.allocQubit(); + std::tie(q0, std::ignore) = b.measure(q0, reg, 0); + auto tensor = b.qtensorFromElements({q0, q1}); + std::tie(tensor, q0) = b.qtensorExtract(tensor, 0); + tensor = b.qtensorInsert(q0, tensor, 0); + std::tie(tensor, q1) = b.qtensorExtract(tensor, 1); + q1 = b.x(q1); + tensor = b.qtensorInsert(q1, tensor, 1); + b.qtensorDealloc(tensor); + return reg; + }); + ASSERT_TRUE(mod); + + std::mt19937_64 rng(11); + auto singleDD = std::make_unique(2); + ASSERT_TRUE(succeeded(sample(mainFunc(*mod), *singleDD, 1, rng))); + const auto singleEvolutionLookups = + singleDD->matrixVectorMultiplication.getStats().lookups; + + auto dd = std::make_unique(2); + const auto histogram = sample(mainFunc(*mod), *dd, 64, rng); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(histogram->at("0") + histogram->at("1"), 64U); + EXPECT_EQ(dd->matrixVectorMultiplication.getStats().lookups, + singleEvolutionLookups); + EXPECT_TRUE(dd->getRootSet().empty()); +} + TEST_F(QCODDFunctionalityTest, SampleDefersAllocatedQubitMeasurement) { auto mod = buildModule([](QCOProgramBuilder& b) { auto reg = @@ -2398,6 +2431,21 @@ TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); } +TEST_F(QCODDFunctionalityTest, RejectsQTensorBeyondQubitRange) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main(%qubits: tensor<65537x!qco.qubit>) { + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + auto dd = std::make_unique(dd::Package::MAX_POSSIBLE_QUBITS); + EXPECT_TRUE(failed(buildFunctionality(mainFunc(*mod), *dd))); +} + TEST_F(QCODDFunctionalityTest, QTensorFlowsThroughLoopAndCall) { auto mod = parseSourceString(R"mlir( module { From bb2d273aafc9ce65e29ee5ea4b7d2f0d83e33ed9 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 19:39:08 +0200 Subject: [PATCH 05/12] =?UTF-8?q?=F0=9F=90=9B=20Make=20QCO=20range=20check?= =?UTF-8?q?s=20portable?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp | 2 +- mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 921edac1bd..a6603bef25 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -1932,7 +1932,7 @@ static bool hasOnlyTerminalQuantumUses(Value value, } if (use.get() == insert.getScalar()) { return !tensorSlot && - hasOnlyTerminalQuantumUses(insert.getResult(), *index, func, + hasOnlyTerminalQuantumUses(insert.getResult(), index, func, outputs, visited); } return tensorSlot && *tensorSlot != *index && diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index a56d6c9de4..9a7628dd81 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -2442,7 +2442,7 @@ TEST_F(QCODDFunctionalityTest, RejectsQTensorBeyondQubitRange) { context.get()); ASSERT_TRUE(mod); - auto dd = std::make_unique(dd::Package::MAX_POSSIBLE_QUBITS); + auto dd = std::make_unique(1); EXPECT_TRUE(failed(buildFunctionality(mainFunc(*mod), *dd))); } From a3b9134c2fcd4222078c2a7b04e21b2d9bdb968f Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 20:22:37 +0200 Subject: [PATCH 06/12] =?UTF-8?q?=F0=9F=90=9B=20Reject=20out-of-range=20st?= =?UTF-8?q?atic=20QCO=20qubits?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 via Codex --- mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp | 7 ++++++- .../Dialect/QCO/Utils/test_dd_functionality.cpp | 12 ++++++++++++ 2 files changed, 18 insertions(+), 1 deletion(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index a6603bef25..2808e04651 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -1733,7 +1733,12 @@ static FailureOr prepare(func::FuncOp func, PreparedState prepared; QubitMap& qubits = prepared.qubits; for (StaticOp staticOp : func.getBody().front().getOps()) { - const auto q = static_cast(staticOp.getIndex()); + const auto index = static_cast(staticOp.getIndex()); + if (index >= dd::Package::MAX_POSSIBLE_QUBITS) { + return staticOp.emitError() + << "static qubit index exceeds the supported qubit range"; + } + const auto q = static_cast(index); qubits.bind(staticOp.getQubit(), q); qubits.numQubits = std::max(qubits.numQubits, static_cast(q) + 1); } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 9a7628dd81..b33af036f9 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -516,6 +516,18 @@ TEST_F(QCODDFunctionalityTest, RejectsUnmappedReturnedQubit) { failed(simulate(mainFunc(*mod), dd::makeZeroState(1, *dd), *dd, rng))); } +TEST_F(QCODDFunctionalityTest, RejectsStaticQubitBeyondDDRange) { + auto mod = buildModule([](QCOProgramBuilder& b) { + b.sink(b.staticQubit(dd::Package::MAX_POSSIBLE_QUBITS)); + return b.intConstant(0); + }); + + auto dd = std::make_unique(1); + EXPECT_TRUE(failed(buildFunctionality(mainFunc(*mod), *dd))); + EXPECT_TRUE( + failed(simulate(mainFunc(*mod), dd::makeZeroState(1, *dd), *dd, rng))); +} + TEST_F(QCODDFunctionalityTest, SimulationConsumesInputReference) { auto valid = buildModule([](QCOProgramBuilder& b) { auto q = b.x(b.staticQubit(0)); From 1ecf027a638131d59553f869e730b963340ecb53 Mon Sep 17 00:00:00 2001 From: Lukas Burgholzer Date: Mon, 31 Aug 2026 19:43:17 +0000 Subject: [PATCH 07/12] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Refine=20QCO=20DD=20?= =?UTF-8?q?runtime=20state=20handling?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Preserve exact MLIR attributes for classical runtime values, validate argument bindings during preparation, reject uninitialized memrefs and non-finite parameters, and transfer QTensor state without repeated slot copies. Assisted-by: GPT-5.6 via Codex --- .../mlir/Dialect/QCO/Utils/DDFunctionality.h | 43 +- mlir/lib/Dialect/QCO/Utils/CMakeLists.txt | 5 +- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 437 +++++++++--------- .../QCO/Utils/test_dd_functionality.cpp | 191 +++----- 4 files changed, 313 insertions(+), 363 deletions(-) diff --git a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h index 46925fb7e6..219996f787 100644 --- a/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h +++ b/mlir/include/mlir/Dialect/QCO/Utils/DDFunctionality.h @@ -28,11 +28,11 @@ namespace mlir::qco { /** * @brief Concrete values for symbolic QCO DD inputs. * - * Integer and `f64` attributes bind scalar function arguments. An integer - * attribute bound to a dynamic one-dimensional qtensor argument gives its - * runtime extent. Bindings for other values are rejected. + * Exactly typed integer, index, and `f64` attributes bind scalar entry-function + * arguments. A non-negative index attribute bound to a dynamic one-dimensional + * qtensor argument gives its runtime extent. Other bindings are rejected. */ -using DDBindings = DenseMap; +using DDArgumentBindings = DenseMap; /** * @brief Sequentially build a matrix DD for a static unitary QCO `func.func`. @@ -67,13 +67,13 @@ using DDBindings = DenseMap; * * @param func The QCO function to construct the functionality for * @param dd The DD package to use (must hold at least the function's qubits) - * @param bindings Concrete scalar values and dynamic QTensor extents for entry - * arguments + * @param argumentBindings Concrete scalar values and dynamic QTensor extents + * for entry arguments * @return The matrix DD on success, or failure for unsupported programs */ -FailureOr -buildFunctionality(func::FuncOp func, dd::Package& dd, - const DDBindings& bindings = DDBindings()); +FailureOr buildFunctionality( + func::FuncOp func, dd::Package& dd, + const DDArgumentBindings& argumentBindings = DDArgumentBindings()); /** * @brief Simulate a QCO `func.func` that may contain measurements, resets, and @@ -86,10 +86,10 @@ buildFunctionality(func::FuncOp func, dd::Package& dd, * arithmetic, comparisons, casts, shifts, and `arith.select`). Dynamic quantum * allocation, qtensors, memrefs, CBit registers, loops, regions, and calls are * supported. QTensor sizes and indices must be concrete; dynamic qtensor - * arguments require an extent in @p bindings. A shared 10000-step budget bounds - * loop iterations across nested loops and calls. Multi-block function bodies - * remain unsupported. Allocated wires stay in the returned state after - * deallocation. + * arguments require an extent in @p argumentBindings. A shared 10000-step + * budget bounds loop iterations and executed control-flow regions and calls. + * Multi-block function bodies remain unsupported. Allocated wires stay in the + * returned state after deallocation. * Consumes one reference to @p in regardless of success or failure. * * @pre The containing module has passed MLIR verification and @@ -100,14 +100,15 @@ buildFunctionality(func::FuncOp func, dd::Package& dd, * higher wires are preserved; one reference is consumed * @param dd The DD package to use * @param rng RNG used for collapsing measurements and resets - * @param bindings Concrete scalar values and dynamic QTensor extents for entry - * arguments + * @param argumentBindings Concrete scalar values and dynamic QTensor extents + * for entry arguments * @return The output statevector DD on success, or failure for unsupported * programs */ -FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, - dd::Package& dd, std::mt19937_64& rng, - const DDBindings& bindings = DDBindings()); +FailureOr +simulate(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, + std::mt19937_64& rng, + const DDArgumentBindings& argumentBindings = DDArgumentBindings()); /** * @brief Sample measurement outcomes from a QCO `func.func`. @@ -129,12 +130,12 @@ FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, * @param dd The DD package to use * @param shots Number of shots * @param rng RNG for collapsing measurements and non-collapsing sampling - * @param bindings Concrete scalar values and dynamic QTensor extents for entry - * arguments + * @param argumentBindings Concrete scalar values and dynamic QTensor extents + * for entry arguments * @return Histogram of outcome strings on success, or failure for unsupported * programs */ FailureOr> sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng, - const DDBindings& bindings = DDBindings()); + const DDArgumentBindings& argumentBindings = DDArgumentBindings()); } // namespace mlir::qco diff --git a/mlir/lib/Dialect/QCO/Utils/CMakeLists.txt b/mlir/lib/Dialect/QCO/Utils/CMakeLists.txt index ddccd4eeb8..62edce61ac 100644 --- a/mlir/lib/Dialect/QCO/Utils/CMakeLists.txt +++ b/mlir/lib/Dialect/QCO/Utils/CMakeLists.txt @@ -73,7 +73,10 @@ add_mlir_library( MLIRSCFDialect MQT::CoreDD PRIVATE - MLIRMQTUtils) + MLIRCBitDialect + MLIRMemRefDialect + MLIRMQTUtils + MLIRQTensorDialect) mqt_mlir_target_use_project_options(MLIRQCODDFunctionality) unset(LLVM_REQUIRES_EH) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 2808e04651..8f7857900e 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -39,7 +39,6 @@ #include #include #include -#include #include #include #include @@ -118,20 +117,29 @@ struct QubitMap { /// Physical wires stored at each tensor index; extracted positions are empty. using TensorSlots = SmallVector>; +using TensorState = std::shared_ptr; struct TensorMap { - DenseMap tensors; + DenseMap tensors; - void bind(Value value, TensorSlots slots) { + void bind(Value value, TensorState slots) { tensors[value] = std::move(slots); } - [[nodiscard]] const TensorSlots* lookup(Value value) const { + [[nodiscard]] TensorState lookup(Value value) const { const auto it = tensors.find(value); - return it == tensors.end() ? nullptr : &it->second; + return it == tensors.end() ? nullptr : it->second; } void erase(Value value) { tensors.erase(value); } + + [[nodiscard]] TensorMap clone() const { + TensorMap copy; + for (const auto& [value, slots] : tensors) { + copy.bind(value, std::make_shared(*slots)); + } + return copy; + } }; struct ClassicalEnv { @@ -140,14 +148,14 @@ struct ClassicalEnv { std::optional deferredWire; }; using RegisterState = std::vector; - using Scalar = std::variant; + using MemRefState = SmallVector; - DenseMap values; + DenseMap values; DenseMap deferredMeasurements; /// Shared storage preserves CBit register identity across `func.call`. DenseMap> registers; /// Shared storage preserves caller-visible writes through `func.call`. - DenseMap>> memrefs; + DenseMap> memrefs; LogicalResult bindFrom(Value source, Value dest, Operation* op) { const auto it = values.find(source); @@ -175,6 +183,10 @@ struct WalkState { size_t remainingExecutionSteps = MAX_CONTROL_FLOW_STEPS; DenseSet activeCalls; }; + +using RuntimeValue = std::variant, + std::shared_ptr>; struct LoopRange { llvm::APInt induction, step; size_t trips; @@ -192,23 +204,23 @@ struct SamplingPlan { isa(tensorType.getElementType()); } -template -static FailureOr lookupScalar(Value value, const ClassicalEnv& classical, - Operation* op) { +static FailureOr +lookupAttribute(Value value, const ClassicalEnv& classical, Operation* op) { const auto it = classical.values.find(value); - if (it == classical.values.end() || !std::holds_alternative(it->second)) { + if (it == classical.values.end()) { return op->emitError() << "classical SSA value is not mapped for QCO DD simulation"; } - return std::get(it->second); + return it->second; } static FailureOr resolveDouble(Value value, const ClassicalEnv& classical, Operation* op) { if (const auto it = classical.values.find(value); - it != classical.values.end() && - std::holds_alternative(it->second)) { - return std::get(it->second); + it != classical.values.end()) { + if (const auto floating = dyn_cast(it->second)) { + return floating.getValue().convertToDouble(); + } } if (const auto constant = mqt::valueToDouble(value)) { return *constant; @@ -263,6 +275,10 @@ decodeStandardGate(UnitaryOpInterface unitary, const ClassicalEnv& classical) { if (failed(concrete)) { return failure(); } + if (!std::isfinite(*concrete)) { + return op->emitError() + << "gate parameters must be finite for QCO DD simulation"; + } decoded.params.push_back(static_cast(*concrete)); } return std::optional{std::move(decoded)}; @@ -325,6 +341,10 @@ static LogicalResult applyUnitaryMatrix(UnitaryOpInterface unitary, if (failed(theta)) { return failure(); } + if (!std::isfinite(*theta)) { + return gphase.emitError() + << "global phase must be finite for QCO DD simulation"; + } auto id = dd::Package::makeIdent(); id.w = walk.dd->cn.lookup(std::cos(*theta), std::sin(*theta)); state = walk.dd->applyOperation(id, state); @@ -421,8 +441,8 @@ static LogicalResult validateReturn(func::ReturnOp returnOp, qc::Qubit expected = 0; for (Value value : returnOp.getOperands()) { if (isQTensorType(value.getType())) { - const auto* slots = tensors.lookup(value); - if (slots == nullptr) { + const auto slots = tensors.lookup(value); + if (!slots) { return returnOp.emitError() << "returned qtensor is not mapped for QCO DD simulation"; } @@ -456,152 +476,111 @@ static LogicalResult validateReturn(func::ReturnOp returnOp, return success(); } +[[nodiscard]] static bool isSupportedClassicalType(Type type) { + return isa(type) || type.isF64(); +} + static LogicalResult recordConstant(arith::ConstantOp constant, ClassicalEnv& classical) { - if (auto attr = dyn_cast(constant.getValue())) { - classical.values[constant.getResult()] = attr.getValue(); - } else if (auto attr = dyn_cast(constant.getValue())) { - if (!constant.getType().isF64()) { - return constant.emitError() - << "QCO DD simulation only supports f64 classical values"; - } - classical.values[constant.getResult()] = attr.getValue().convertToDouble(); - } else if (auto attr = dyn_cast(constant.getValue())) { - if (constant.getType().isInteger(1)) { - classical.values[constant.getResult()] = attr.getValue() != 0; - } else if (isa(constant.getType())) { - classical.values[constant.getResult()] = attr.getInt(); - } else if (isa(constant.getType())) { - classical.values[constant.getResult()] = attr.getValue(); - } + if (!isSupportedClassicalType(constant.getType())) { + return constant.emitError() + << "QCO DD simulation only supports integer, index, and f64 values"; } + classical.values[constant.getResult()] = constant.getValue(); return success(); } -static LogicalResult applyBindings(func::FuncOp func, - const DDBindings& bindings, - ClassicalEnv& classical) { - for (const auto& [value, attr] : bindings) { - auto argument = dyn_cast(value); - if (!argument || argument.getOwner() != &func.getBody().front()) { - return func.emitError() - << "QCO DD bindings must target entry-block arguments"; +static LogicalResult +applyArgumentBindings(func::FuncOp func, + const DDArgumentBindings& argumentBindings, + ClassicalEnv& classical) { + size_t boundArguments = 0; + for (Value argument : func.getArguments()) { + const auto binding = argumentBindings.find(argument); + if (binding == argumentBindings.end()) { + continue; } - const Type type = value.getType(); + ++boundArguments; + const Attribute attr = binding->second; + const Type type = argument.getType(); if (isQTensorType(type)) { - if (cast(type).isDynamicDim(0) && - isa(attr)) { - continue; - } - } else if (type.isInteger(1)) { - if (auto boolean = dyn_cast(attr)) { - classical.values[value] = boolean.getValue(); - continue; - } - if (auto integer = dyn_cast(attr)) { - classical.values[value] = integer.getValue() != 0; - continue; - } - } else if (isa(type)) { - if (auto integer = dyn_cast(attr)) { - classical.values[value] = integer.getInt(); - continue; - } - } else if (auto integerType = dyn_cast(type)) { - if (auto integer = dyn_cast(attr)) { - classical.values[value] = - integer.getValue().sextOrTrunc(integerType.getWidth()); - continue; - } - } else if (type.isF64()) { - if (auto floating = dyn_cast(attr); - floating && floating.getType().isF64()) { - classical.values[value] = floating.getValue().convertToDouble(); + const auto extent = dyn_cast(attr); + if (cast(type).isDynamicDim(0) && extent && + isa(extent.getType())) { continue; } + } else if (const auto typed = dyn_cast(attr); + isSupportedClassicalType(type) && typed && + typed.getType() == type) { + classical.values[argument] = attr; + continue; } return func.emitError() << "QCO DD binding attribute " << attr << " does not match argument type " << type; } + if (boundArguments != argumentBindings.size()) { + return func.emitError() + << "QCO DD bindings must target entry-block arguments"; + } return success(); } static FailureOr lookupBool(Value value, const ClassicalEnv& classical, Operation* op) { - return lookupScalar(value, classical, op); + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer || !value.getType().isInteger(1)) { + return failure(); + } + return !integer.getValue().isZero(); } static FailureOr lookupIndex(Value value, const ClassicalEnv& classical, Operation* op) { - return lookupScalar(value, classical, op); + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer || !isa(value.getType())) { + return failure(); + } + return integer.getValue().getSExtValue(); } static FailureOr lookupFloat(Value value, const ClassicalEnv& classical, Operation* op) { - return lookupScalar(value, classical, op); + auto attr = lookupAttribute(value, classical, op); + const auto floating = + succeeded(attr) ? dyn_cast(*attr) : FloatAttr{}; + if (!floating || !value.getType().isF64()) { + return failure(); + } + return floating.getValue().convertToDouble(); } static FailureOr lookupInteger(Value value, const ClassicalEnv& classical, Operation* op) { - if (value.getType().isInteger(1)) { - auto bit = lookupBool(value, classical, op); - if (failed(bit)) { - return failure(); - } - return llvm::APInt(1, static_cast(*bit)); + if (!isa(value.getType())) { + return op->emitError() << "expected an integer or index SSA value"; } - if (isa(value.getType())) { - auto index = lookupIndex(value, classical, op); - if (failed(index)) { - return failure(); - } - return llvm::APInt(64, static_cast(*index)); - } - if (isa(value.getType())) { - return lookupScalar(value, classical, op); + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer) { + return failure(); } - return op->emitError() << "expected an integer or index SSA value"; -} - -[[nodiscard]] static bool evaluateCmp(arith::CmpIPredicate predicate, - const llvm::APInt& lhs, - const llvm::APInt& rhs) { - switch (predicate) { - case arith::CmpIPredicate::eq: - return lhs == rhs; - case arith::CmpIPredicate::ne: - return lhs != rhs; - case arith::CmpIPredicate::slt: - return lhs.slt(rhs); - case arith::CmpIPredicate::sle: - return lhs.sle(rhs); - case arith::CmpIPredicate::sgt: - return lhs.sgt(rhs); - case arith::CmpIPredicate::sge: - return lhs.sge(rhs); - case arith::CmpIPredicate::ult: - return lhs.ult(rhs); - case arith::CmpIPredicate::ule: - return lhs.ule(rhs); - case arith::CmpIPredicate::ugt: - return lhs.ugt(rhs); - case arith::CmpIPredicate::uge: - return lhs.uge(rhs); - } - llvm_unreachable("unknown arith.cmpi predicate"); + return integer.getValue(); } static LogicalResult bindInteger(Value dest, const llvm::APInt& value, ClassicalEnv& classical) { - if (dest.getType().isInteger(1)) { - classical.values[dest] = value[0]; - } else if (isa(dest.getType())) { - classical.values[dest] = static_cast(value.getZExtValue()); - } else if (auto type = dyn_cast(dest.getType())) { - classical.values[dest] = value.zextOrTrunc(type.getWidth()); - } else { + const Type type = dest.getType(); + if (!isa(type)) { return failure(); } + const unsigned width = + isa(type) ? 64U : cast(type).getWidth(); + classical.values[dest] = IntegerAttr::get(type, value.zextOrTrunc(width)); return success(); } @@ -682,13 +661,9 @@ static LogicalResult loadRegister(cbit::LoadOp load, ClassicalEnv& classical) { classical); } -[[nodiscard]] static bool isSupportedClassicalType(Type type) { - return isa(type) || type.isF64(); -} - -static FailureOr -lookupMemRefSlot(Value memref, ValueRange indices, ClassicalEnv& classical, - Operation* op) { +static FailureOr lookupMemRefSlot(Value memref, ValueRange indices, + ClassicalEnv& classical, + Operation* op) { const auto type = dyn_cast(memref.getType()); if (!type || type.getRank() != 1 || indices.size() != 1 || !isSupportedClassicalType(type.getElementType())) { @@ -712,19 +687,6 @@ lookupMemRefSlot(Value memref, ValueRange indices, ClassicalEnv& classical, return &(*it->second)[static_cast(*index)]; } -static ClassicalEnv::Scalar zeroScalar(Type type) { - if (type.isInteger(1)) { - return false; - } - if (isa(type)) { - return int64_t{0}; - } - if (auto integer = dyn_cast(type)) { - return llvm::APInt(integer.getWidth(), 0); - } - return 0.0; -} - static LogicalResult applyMemRefAlloc(memref::AllocOp alloc, ClassicalEnv& classical) { const auto type = dyn_cast(alloc.getType()); @@ -754,8 +716,7 @@ static LogicalResult applyMemRefAlloc(memref::AllocOp alloc, return alloc.emitError() << "classical memref size must be non-negative"; } classical.memrefs[alloc.getResult()] = - std::make_shared>( - static_cast(size), zeroScalar(type.getElementType())); + std::make_shared(static_cast(size)); return success(); } @@ -782,6 +743,10 @@ static LogicalResult applyMemRefLoad(memref::LoadOp load, if (failed(slot)) { return failure(); } + if (!**slot) { + return load.emitError() + << "read from an uninitialized classical memref element"; + } classical.values[load.getResult()] = **slot; return success(); } @@ -805,7 +770,8 @@ static LogicalResult applyBinaryFloat(OpTy op, ClassicalEnv& classical, if (failed(lhs) || failed(rhs)) { return failure(); } - classical.values[op.getResult()] = combine(*lhs, *rhs); + classical.values[op.getResult()] = + FloatAttr::get(op.getResult().getType(), combine(*lhs, *rhs)); return success(); } @@ -958,8 +924,9 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { if (failed(lhs) || failed(rhs)) { return failure(); } - classical.values[cmp.getResult()] = - evaluateCmp(cmp.getPredicate(), *lhs, *rhs); + classical.values[cmp.getResult()] = BoolAttr::get( + cmp.getContext(), + arith::applyCmpPredicate(cmp.getPredicate(), *lhs, *rhs)); return success(); }) .Case([&](arith::SelectOp select) -> LogicalResult { @@ -1017,7 +984,8 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { if (failed(value)) { return failure(); } - classical.values[neg.getResult()] = -*value; + classical.values[neg.getResult()] = + FloatAttr::get(neg.getType(), -*value); return success(); }) .Case([&](arith::CmpFOp cmp) -> LogicalResult { @@ -1026,8 +994,10 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { if (failed(lhs) || failed(rhs)) { return failure(); } - classical.values[cmp.getResult()] = arith::applyCmpPredicate( - cmp.getPredicate(), llvm::APFloat(*lhs), llvm::APFloat(*rhs)); + classical.values[cmp.getResult()] = BoolAttr::get( + cmp.getContext(), + arith::applyCmpPredicate(cmp.getPredicate(), llvm::APFloat(*lhs), + llvm::APFloat(*rhs))); return success(); }) .Case( @@ -1037,8 +1007,9 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { if (failed(value)) { return failure(); } - classical.values[castOp->getResult(0)] = - value->roundToDouble(isa(castOp)); + classical.values[castOp->getResult(0)] = FloatAttr::get( + castOp->getResult(0).getType(), + value->roundToDouble(isa(castOp))); return success(); }) .Case( @@ -1120,45 +1091,60 @@ static LogicalResult bindLinearArgs(ValueRange operands, Block& block, static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, WalkState& walk, Operation* op) { - const QubitMap sourceQubits = *walk.qubits; - const TensorMap sourceTensors = *walk.tensors; - const ClassicalEnv sourceClassical = *walk.classical; + SmallVector values; + values.reserve(sources.size()); for (auto [src, dest] : llvm::zip_equal(sources, dests)) { if (isa(dest.getType())) { - const auto q = sourceQubits.lookup(src); + const auto q = walk.qubits->lookup(src); if (!q) { return op->emitError() << "qubit SSA value is not mapped for QCO DD construction"; } - walk.qubits->bind(dest, *q); + values.emplace_back(*q); } else if (isQTensorType(dest.getType())) { - const auto* slots = sourceTensors.lookup(src); - if (slots == nullptr) { + const auto slots = walk.tensors->lookup(src); + if (!slots) { return op->emitError() << "qtensor SSA value is not mapped for QCO DD simulation"; } - walk.tensors->bind(dest, *slots); + values.emplace_back(slots); } else if (isa(dest.getType())) { - const auto it = sourceClassical.registers.find(src); - if (it == sourceClassical.registers.end()) { + const auto it = walk.classical->registers.find(src); + if (it == walk.classical->registers.end()) { return op->emitError() << "CBit register is not mapped for QCO DD simulation"; } - walk.classical->registers[dest] = it->second; + values.emplace_back(it->second); } else if (isa(dest.getType())) { - const auto it = sourceClassical.memrefs.find(src); - if (it == sourceClassical.memrefs.end()) { + const auto it = walk.classical->memrefs.find(src); + if (it == walk.classical->memrefs.end()) { return op->emitError() << "classical memref is not mapped for QCO DD simulation"; } - walk.classical->memrefs[dest] = it->second; + values.emplace_back(it->second); } else { - const auto value = sourceClassical.values.find(src); - if (value == sourceClassical.values.end()) { + const auto value = walk.classical->values.find(src); + if (value == walk.classical->values.end()) { return op->emitError() << "classical SSA value is not mapped for QCO DD simulation"; } - walk.classical->values[dest] = value->second; + values.emplace_back(value->second); + } + } + + for (auto [value, dest] : llvm::zip_equal(values, dests)) { + if (isa(dest.getType())) { + walk.qubits->bind(dest, std::get(value)); + } else if (isQTensorType(dest.getType())) { + walk.tensors->bind(dest, std::get(value)); + } else if (isa(dest.getType())) { + walk.classical->registers[dest] = + std::get>(value); + } else if (isa(dest.getType())) { + walk.classical->memrefs[dest] = + std::get>(value); + } else { + walk.classical->values[dest] = std::get(value); } } return success(); @@ -1303,7 +1289,9 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { if (failed(slots)) { return failure(); } - walk.tensors->bind(alloc.getResult(), std::move(*slots)); + walk.tensors->bind( + alloc.getResult(), + std::make_shared(std::move(*slots))); return success(); } }) @@ -1319,16 +1307,17 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { for (const qc::Qubit wire : *wires) { slots.emplace_back(wire); } - walk.tensors->bind(fromElements.getResult(), std::move(slots)); + walk.tensors->bind(fromElements.getResult(), + std::make_shared(std::move(slots))); return success(); }) .template Case( [&](qtensor::ExtractOp extract) -> LogicalResult { - const auto* input = walk.tensors->lookup(extract.getTensor()); + const auto input = walk.tensors->lookup(extract.getTensor()); auto index = lookupIndex(extract.getIndex(), *walk.classical, extract); - if (input == nullptr || failed(index)) { - if (input == nullptr) { + if (!input || failed(index)) { + if (!input) { extract.emitError() << "qtensor is not mapped for QCO DD simulation"; } @@ -1337,25 +1326,24 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { if (*index < 0 || static_cast(*index) >= input->size()) { return extract.emitError() << "qtensor index out of range"; } - TensorSlots output = *input; - auto& wire = output[static_cast(*index)]; + auto& wire = (*input)[static_cast(*index)]; if (!wire) { return extract.emitError() << "qtensor element has already been extracted"; } walk.qubits->bind(extract.getResult(), *wire); wire.reset(); - walk.tensors->bind(extract.getOutTensor(), std::move(output)); + walk.tensors->bind(extract.getOutTensor(), input); return success(); }) .template Case( [&](qtensor::InsertOp insert) -> LogicalResult { - const auto* input = walk.tensors->lookup(insert.getDest()); + const auto input = walk.tensors->lookup(insert.getDest()); const auto wire = walk.qubits->lookup(insert.getScalar()); auto index = lookupIndex(insert.getIndex(), *walk.classical, insert); - if (input == nullptr || !wire || failed(index)) { - if (input == nullptr || !wire) { + if (!input || !wire || failed(index)) { + if (!input || !wire) { insert.emitError() << "qtensor or qubit is not mapped for QCO DD simulation"; } @@ -1364,14 +1352,13 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { if (*index < 0 || static_cast(*index) >= input->size()) { return insert.emitError() << "qtensor index out of range"; } - TensorSlots output = *input; - output[static_cast(*index)] = wire; - walk.tensors->bind(insert.getResult(), std::move(output)); + (*input)[static_cast(*index)] = wire; + walk.tensors->bind(insert.getResult(), input); return success(); }) .template Case( [&](qtensor::DeallocOp dealloc) -> LogicalResult { - if (walk.tensors->lookup(dealloc.getTensor()) == nullptr) { + if (!walk.tensors->lookup(dealloc.getTensor())) { return dealloc.emitError() << "qtensor is not mapped for QCO DD simulation"; } @@ -1436,7 +1423,8 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return success(); } const char bit = walk.dd->measureOneCollapsing(state, *q, *walk.rng); - walk.classical->values[measureOp.getResult()] = bit == '1'; + walk.classical->values[measureOp.getResult()] = + BoolAttr::get(measureOp.getContext(), bit == '1'); walk.qubits->bind(measureOp.getQubitOut(), *q); return success(); } @@ -1718,19 +1706,24 @@ namespace { struct PreparedState { QubitMap qubits; TensorMap tensors; + ClassicalEnv classical; }; } // namespace -static FailureOr prepare(func::FuncOp func, - const dd::Package& dd, - const DDBindings& bindings, - bool bindEntryAllocations = false) { +static FailureOr +prepare(func::FuncOp func, const dd::Package& dd, + const DDArgumentBindings& argumentBindings, + bool bindEntryAllocations = false) { if (!func.getBody().hasOneBlock()) { return func.emitError() << "QCO DD construction expects a single-block function body"; } PreparedState prepared; + if (failed( + applyArgumentBindings(func, argumentBindings, prepared.classical))) { + return failure(); + } QubitMap& qubits = prepared.qubits; for (StaticOp staticOp : func.getBody().front().getOps()) { const auto index = static_cast(staticOp.getIndex()); @@ -1755,12 +1748,12 @@ static FailureOr prepare(func::FuncOp func, const auto type = cast(arg.getType()); int64_t size = type.getDimSize(0); if (type.isDynamicDim(0)) { - const auto binding = bindings.find(arg); - if (binding == bindings.end() || !isa(binding->second)) { + const auto binding = argumentBindings.find(arg); + if (binding == argumentBindings.end()) { return func.emitError() - << "dynamic qtensor arguments require an integer extent"; + << "dynamic qtensor arguments require an index extent"; } - size = cast(binding->second).getInt(); + size = cast(binding->second).getValue().getSExtValue(); if (size < 0) { return func.emitError() << "dynamic qtensor extent must be non-negative"; @@ -1776,13 +1769,18 @@ static FailureOr prepare(func::FuncOp func, for (size_t i = 0; i < count; ++i) { slots.emplace_back(static_cast(next++)); } - prepared.tensors.bind(arg, std::move(slots)); + prepared.tensors.bind(arg, + std::make_shared(std::move(slots))); } } qubits.numQubits = next; } if (bindEntryAllocations) { for (AllocOp alloc : func.getBody().front().getOps()) { + if (qubits.numQubits >= dd::Package::MAX_POSSIBLE_QUBITS) { + return alloc.emitError() + << "QCO function exceeds the supported qubit range"; + } qubits.bind(alloc.getResult(), static_cast(qubits.numQubits++)); } @@ -1794,18 +1792,17 @@ static FailureOr prepare(func::FuncOp func, return prepared; } -FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd, - const DDBindings& bindings) { - auto prepared = prepare(func, dd, bindings, /*bindEntryAllocations=*/true); +FailureOr +buildFunctionality(func::FuncOp func, dd::Package& dd, + const DDArgumentBindings& argumentBindings) { + auto prepared = + prepare(func, dd, argumentBindings, /*bindEntryAllocations=*/true); if (failed(prepared)) { return failure(); } QubitMap qubits = std::move(prepared->qubits); TensorMap tensors = std::move(prepared->tensors); - ClassicalEnv classical; - if (failed(applyBindings(func, bindings, classical))) { - return failure(); - } + ClassicalEnv classical = std::move(prepared->classical); WalkState walkState{.qubits = &qubits, .tensors = &tensors, .classical = &classical, @@ -1828,7 +1825,6 @@ FailureOr buildFunctionality(func::FuncOp func, dd::Package& dd, static FailureOr simulateImpl(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, const PreparedState& prepared, std::mt19937_64* rng, - const DDBindings& bindings, const DenseSet* deferredMeasurements = nullptr, ClassicalEnv* finalClassical = nullptr) { const size_t inputQubits = @@ -1841,12 +1837,8 @@ simulateImpl(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, } QubitMap qubits = prepared.qubits; qubits.numQubits = inputQubits; - TensorMap tensors = prepared.tensors; - ClassicalEnv classical; - if (failed(applyBindings(func, bindings, classical))) { - dd.decRef(in); - return failure(); - } + TensorMap tensors = prepared.tensors.clone(); + ClassicalEnv classical = prepared.classical; WalkState walkState{.qubits = &qubits, .tensors = &tensors, .classical = &classical, @@ -1867,13 +1859,13 @@ simulateImpl(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, FailureOr simulate(func::FuncOp func, const dd::VectorDD& in, dd::Package& dd, std::mt19937_64& rng, - const DDBindings& bindings) { - auto prepared = prepare(func, dd, bindings); + const DDArgumentBindings& argumentBindings) { + auto prepared = prepare(func, dd, argumentBindings); if (failed(prepared)) { dd.decRef(in); return failure(); } - return simulateImpl(func, in, dd, *prepared, &rng, bindings); + return simulateImpl(func, in, dd, *prepared, &rng); } static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { @@ -2052,11 +2044,10 @@ static FailureOr encodeOutcome(ArrayRef outputs, return outcome; } -FailureOr> sample(func::FuncOp func, - dd::Package& dd, size_t shots, - std::mt19937_64& rng, - const DDBindings& bindings) { - auto prepared = prepare(func, dd, bindings); +FailureOr> +sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng, + const DDArgumentBindings& argumentBindings) { + auto prepared = prepare(func, dd, argumentBindings); if (failed(prepared)) { return failure(); } @@ -2083,9 +2074,9 @@ FailureOr> sample(func::FuncOp func, if (!plan->dynamic) { ClassicalEnv classical; - auto state = simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, - *prepared, nullptr, bindings, - &plan->deferredMeasurements, &classical); + auto state = + simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, *prepared, + nullptr, &plan->deferredMeasurements, &classical); if (failed(state)) { return failure(); } @@ -2101,7 +2092,7 @@ FailureOr> sample(func::FuncOp func, for (size_t i = 0; i < shots; ++i) { ClassicalEnv classical; auto state = simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, - *prepared, &rng, bindings, nullptr, &classical); + *prepared, &rng, nullptr, &classical); if (failed(state)) { return failure(); } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index b33af036f9..5cf49a7d10 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -1722,7 +1722,7 @@ TEST_F(QCODDFunctionalityTest, RejectsUnsupportedFuncCalls) { context.get()); ASSERT_TRUE(selfRecursive); auto selfRecursiveFunc = mainFunc(*selfRecursive); - DDBindings bindings; + DDArgumentBindings bindings; bindings[selfRecursiveFunc.getArgument(0)] = BoolAttr::get(context.get(), true); auto zeroQubitDd = std::make_unique(0); @@ -2242,7 +2242,7 @@ TEST_F(QCODDFunctionalityTest, SymbolicParametersUseBindings) { ASSERT_TRUE(concrete); auto func = mainFunc(*mod); - DDBindings bindings; + DDArgumentBindings bindings; bindings[func.getArgument(0)] = FloatAttr::get( cast(func.getArgument(0).getType()), std::numbers::pi / 2.0); @@ -2269,6 +2269,43 @@ TEST_F(QCODDFunctionalityTest, SymbolicParametersUseBindings) { EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); } +TEST_F(QCODDFunctionalityTest, RejectsNonFiniteParameters) { + auto gate = parseSourceString(R"mlir( + module { + func.func @main(%theta: f64) { + %q = qco.static 0 : !qco.qubit + %out = qco.rx(%theta) %q : !qco.qubit -> !qco.qubit + qco.sink %out : !qco.qubit + return + } + } + )mlir", + context.get()); + auto phase = parseSourceString(R"mlir( + module { + func.func @main(%theta: f64) { + qco.gphase(%theta) + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(gate); + ASSERT_TRUE(phase); + + for (auto [func, value] : + {std::pair{mainFunc(*gate), std::numeric_limits::infinity()}, + std::pair{mainFunc(*phase), + std::numeric_limits::quiet_NaN()}}) { + DDArgumentBindings bindings; + bindings[func.getArgument(0)] = + FloatAttr::get(Float64Type::get(context.get()), value); + auto dd = std::make_unique(1); + EXPECT_TRUE(failed(sample(func, *dd, 1, rng, bindings))); + EXPECT_TRUE(dd->getRootSet().empty()); + } +} + TEST_F(QCODDFunctionalityTest, BuildsThroughConcreteControlFlow) { auto mod = buildModule([](QCOProgramBuilder& b) { auto q = b.staticQubit(0); @@ -2427,7 +2464,7 @@ TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { context.get()); ASSERT_TRUE(mod); auto func = mainFunc(*mod); - DDBindings bindings; + DDArgumentBindings bindings; bindings[func.getArgument(0)] = IntegerAttr::get(IndexType::get(context.get()), 2); @@ -2441,6 +2478,9 @@ TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { bindings[func.getArgument(0)] = IntegerAttr::get(IndexType::get(context.get()), -1); EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); + bindings[func.getArgument(0)] = + IntegerAttr::get(IntegerType::get(context.get(), 64), 2); + EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); } TEST_F(QCODDFunctionalityTest, RejectsQTensorBeyondQubitRange) { @@ -2553,44 +2593,26 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { expectSimulatesFromZero(mainFunc(*mod), true); } -TEST_F(QCODDFunctionalityTest, - SupportsAdditionalClassicalOperationsAndBindings) { +TEST_F(QCODDFunctionalityTest, BindingsDriveObservableClassicalPath) { auto mod = parseSourceString(R"mlir( module { func.func @main(%idx: index, %word: i16, %flag: i1) { - %zero = arith.constant 0 : i8 - %one = arith.constant 1 : i8 - %two = arith.constant 2 : i8 - %four = arith.constant 4 : i8 - %negative = arith.constant -5 : i8 - %sle = arith.cmpi sle, %one, %two : i8 - %sgt = arith.cmpi sgt, %two, %one : i8 - %sge = arith.cmpi sge, %two, %two : i8 - %ult = arith.cmpi ult, %one, %two : i8 - %ule = arith.cmpi ule, %two, %two : i8 - %ugt = arith.cmpi ugt, %two, %one : i8 - %uge = arith.cmpi uge, %two, %two : i8 - %quotient = arith.divui %four, %two : i8 - %remainder = arith.remsi %negative, %two : i8 - %shifted = arith.shrsi %negative, %one : i8 - %extended = arith.extsi %negative : i8 to i16 + %one = arith.constant 1 : i16 + %expected = arith.constant 4 : i16 + %sum = arith.addi %word, %one : i16 + %integer_ok = arith.cmpi eq, %sum, %expected : i16 %as_index = arith.index_cast %word : i16 to index - %selected = arith.select %flag, %one, %zero : i8 - %one_float = arith.constant 1.0 : f64 - %two_float = arith.constant 2.0 : f64 - %difference = arith.subf %two_float, %one_float : f64 - %product = arith.mulf %difference, %two_float : f64 - %negated = arith.negf %product : f64 - %zero_index = arith.constant 0 : index - %indices = memref.alloc() : memref<1xindex> - %loaded_index = memref.load %indices[%zero_index] : memref<1xindex> - memref.dealloc %indices : memref<1xindex> - %floats = memref.alloc() : memref<1xf64> - %loaded_float = memref.load %floats[%zero_index] : memref<1xf64> - memref.dealloc %floats : memref<1xf64> - %false = arith.constant false - scf.if %false { + %index_ok = arith.cmpi eq, %as_index, %idx : index + %both = arith.andi %integer_ok, %index_ok : i1 + %condition = arith.select %flag, %both, %flag : i1 + %q = qco.static 0 : !qco.qubit + %out = qco.if %condition args(%arg = %q) -> (!qco.qubit) { + %flipped = qco.x %arg : !qco.qubit -> !qco.qubit + qco.yield %flipped : !qco.qubit + } else args(%arg = %q) { + qco.yield %arg : !qco.qubit } + qco.sink %out : !qco.qubit return } } @@ -2599,21 +2621,25 @@ TEST_F(QCODDFunctionalityTest, ASSERT_TRUE(mod); auto func = mainFunc(*mod); - DDBindings bindings; + DDArgumentBindings bindings; bindings[func.getArgument(0)] = IntegerAttr::get(IndexType::get(context.get()), 3); bindings[func.getArgument(1)] = - IntegerAttr::get(IntegerType::get(context.get(), 16), -2); + IntegerAttr::get(IntegerType::get(context.get(), 16), 3); bindings[func.getArgument(2)] = IntegerAttr::get(IntegerType::get(context.get(), 1), 1); - auto dd = std::make_unique(0); - const auto output = simulate(func, dd::VectorDD::one(), *dd, rng, bindings); - ASSERT_TRUE(succeeded(output)); - EXPECT_TRUE(output->isTerminal()); - dd->decRef(*output); + auto dd = std::make_unique(1); + const auto histogram = sample(func, *dd, 4, rng, bindings); + ASSERT_TRUE(succeeded(histogram)); + EXPECT_EQ(*histogram, (std::map{{"1", 4}})); + + bindings[func.getArgument(1)] = + IntegerAttr::get(IntegerType::get(context.get(), 8), 3); + EXPECT_TRUE(failed(sample(func, *dd, 1, rng, bindings))); + EXPECT_TRUE(dd->getRootSet().empty()); } -TEST_F(QCODDFunctionalityTest, RejectsClassicalRuntimeErrors) { +TEST_F(QCODDFunctionalityTest, RejectsRepresentativeClassicalRuntimeErrors) { for (const StringRef source : { R"mlir(module { func.func @main(%unbound: i16) { @@ -2631,20 +2657,13 @@ TEST_F(QCODDFunctionalityTest, RejectsClassicalRuntimeErrors) { } })mlir", R"mlir(module { - func.func @main(%index: index) { + func.func @main() { + %index = arith.constant 0 : index %reg = memref.alloc() : memref<1xi16> %value = memref.load %reg[%index] : memref<1xi16> return } })mlir", - R"mlir(module { - func.func @main(%value: i16) { - %zero = arith.constant 0 : index - %reg = memref.alloc() : memref<1xi16> - memref.store %value, %reg[%zero] : memref<1xi16> - return - } - })mlir", R"mlir(module { func.func @main() { %negative = arith.constant -1 : index @@ -2667,70 +2686,6 @@ TEST_F(QCODDFunctionalityTest, RejectsClassicalRuntimeErrors) { return } })mlir", - R"mlir(module { - func.func @main(%rhs: i8) { - %one = arith.constant 1 : i8 - %invalid = arith.divui %one, %rhs : i8 - return - } - })mlir", - R"mlir(module { - func.func @main(%lhs: i8) { - %one = arith.constant 1 : i8 - %invalid = arith.divui %lhs, %one : i8 - return - } - })mlir", - R"mlir(module { - func.func @main(%amount: i8) { - %one = arith.constant 1 : i8 - %invalid = arith.shli %one, %amount : i8 - return - } - })mlir", - R"mlir(module { - func.func @main(%lhs: f64) { - %zero = arith.constant 0.0 : f64 - %invalid = arith.cmpf oeq, %lhs, %zero : f64 - return - } - })mlir", - R"mlir(module { - func.func @main(%value: i8) { - %invalid = arith.sitofp %value : i8 to f64 - return - } - })mlir", - R"mlir(module { - func.func @main(%value: f64) { - %invalid = arith.fptosi %value : f64 to i8 - return - } - })mlir", - R"mlir(module { - func.func @consume(%reg: memref<1xi16>) { - return - } - func.func @main(%reg: memref<1xi16>) { - func.call @consume(%reg) : (memref<1xi16>) -> () - return - } - })mlir", - R"mlir(module { - func.func @main(%condition: i1) { - scf.if %condition { - } - return - } - })mlir", - R"mlir(module { - func.func @main(%selector: index) { - scf.index_switch %selector - default { - } - return - } - })mlir", R"mlir(module { func.func @main(%size: index) { %tensor = qtensor.alloc(%size) : tensor @@ -2753,7 +2708,7 @@ TEST_F(QCODDFunctionalityTest, RejectsClassicalRuntimeErrors) { ASSERT_TRUE(mod); auto func = mainFunc(*mod); auto constant = *func.getBody().front().getOps().begin(); - DDBindings bindings; + DDArgumentBindings bindings; bindings[constant.getResult()] = IntegerAttr::get(IndexType::get(context.get()), 0); auto dd = std::make_unique(0); From 618d38a1e70620a424388f68b378f04ef4c8bf24 Mon Sep 17 00:00:00 2001 From: Lukas Burgholzer Date: Mon, 31 Aug 2026 19:43:38 +0000 Subject: [PATCH 08/12] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Simplify=20QCO=20DD?= =?UTF-8?q?=20sampling=20and=20execution=20guards?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Restore the conservative entry-block measurement classifier, centralize the shared execution budget, and remove speculative QTensor terminal-use analysis and repetitive coverage. Assisted-by: GPT-5.6 via Codex --- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 134 ++++++------------ .../QCO/Utils/test_dd_functionality.cpp | 66 ++++----- 2 files changed, 73 insertions(+), 127 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 8f7857900e..5687bddfba 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -198,6 +198,15 @@ struct SamplingPlan { }; } // namespace +static LogicalResult consumeExecutionStep(WalkState& walk, Operation* op) { + if (walk.remainingExecutionSteps == 0) { + return op->emitError( + "QCO DD execution exceeds the limit of 10000 control-flow steps"); + } + --walk.remainingExecutionSteps; + return success(); +} + [[nodiscard]] static bool isQTensorType(Type type) { const auto tensorType = dyn_cast(type); return tensorType && tensorType.getRank() == 1 && @@ -1039,8 +1048,8 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { }); } -static FailureOr -resolveLoop(scf::ForOp forOp, ClassicalEnv& classical, size_t remainingSteps) { +static FailureOr resolveLoop(scf::ForOp forOp, + ClassicalEnv& classical) { auto lower = lookupInteger(forOp.getLowerBound(), classical, forOp); auto upper = lookupInteger(forOp.getUpperBound(), classical, forOp); auto step = lookupInteger(forOp.getStep(), classical, forOp); @@ -1067,11 +1076,7 @@ resolveLoop(scf::ForOp forOp, ClassicalEnv& classical, size_t remainingSteps) { const llvm::APInt span = upperWide - lowerWide; const llvm::APInt trips = (span + stepWide - llvm::APInt(wideWidth, 1)).udiv(stepWide); - const size_t limited = trips.getLimitedValue(remainingSteps + 1); - if (limited > remainingSteps) { - return forOp.emitError( - "QCO DD execution exceeds the limit of 10000 control-flow steps"); - } + const size_t limited = trips.getLimitedValue(MAX_CONTROL_FLOW_STEPS + 1); return LoopRange{.induction = lowerWide, .step = stepWide, .trips = limited}; } @@ -1188,6 +1193,9 @@ applyRegionBranch(ValueRange linearOperands, Block& block, if (failed(bindLinearArgs(linearOperands, block, walk, parent))) { return failure(); } + if (failed(consumeExecutionStep(walk, parent))) { + return failure(); + } if (failed(walkBlock(block, walk, state))) { return failure(); } @@ -1204,6 +1212,9 @@ static LogicalResult applyScfRegion(Region& region, ValueRange results, << "SCF region must contain exactly one block for QCO DD simulation"; } Block& block = region.front(); + if (failed(consumeExecutionStep(walk, parent))) { + return failure(); + } if (failed(walkBlock(block, walk, state))) { return failure(); } @@ -1527,8 +1538,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { walk, state, execute); }) .template Case([&](scf::ForOp forOp) -> LogicalResult { - auto range = - resolveLoop(forOp, *walk.classical, walk.remainingExecutionSteps); + auto range = resolveLoop(forOp, *walk.classical); if (failed(range)) { return failure(); } @@ -1539,12 +1549,9 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { for (size_t t = 0; t < range->trips; ++t, range->induction += range->step) { - if (walk.remainingExecutionSteps == 0) { - return forOp.emitError( - "QCO DD execution exceeds the limit of 10000 control-flow " - "steps"); + if (failed(consumeExecutionStep(walk, forOp))) { + return failure(); } - --walk.remainingExecutionSteps; auto iterArgs = body.getArguments().drop_front(); if (failed(bindValuePairs(carried, iterArgs, walk, forOp))) { return failure(); @@ -1594,12 +1601,9 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return bindValuePairs(condition.getArgs(), whileOp.getResults(), walk, whileOp); } - if (walk.remainingExecutionSteps == 0) { - return whileOp.emitError( - "QCO DD execution exceeds the limit of 10000 control-flow " - "steps"); + if (failed(consumeExecutionStep(walk, whileOp))) { + return failure(); } - --walk.remainingExecutionSteps; if (failed(bindValuePairs(condition.getArgs(), after.getArguments(), walk, whileOp)) || failed(walkBlock(after, walk, state))) { @@ -1632,6 +1636,10 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { const auto guard = llvm::make_scope_exit([&] { walk.activeCalls.erase(calleeOp); }); + if (failed(consumeExecutionStep(walk, call))) { + return failure(); + } + if (failed(bindValuePairs(call.getArgOperands(), callee.getArguments(), walk, call))) { return failure(); @@ -1880,89 +1888,27 @@ static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { static bool hasOutputOnlyMeasurementResult(MeasureOp measure, ArrayRef outputs) { - if (measure.getResult().use_empty()) { - return false; - } return llvm::all_of(measure.getResult().getUses(), [&](const OpOperand& use) { auto store = dyn_cast(use.getOwner()); return store && isOutputOnlyRegister(store.getReg(), outputs); }); } -static std::optional getConstantTensorIndex(Value value) { - const auto attr = mqt::valueToConstantAttr(value); - const auto integer = attr ? dyn_cast(*attr) : IntegerAttr{}; - if (!integer || integer.getInt() < 0) { - return std::nullopt; - } - return static_cast(integer.getInt()); -} - -static bool hasOnlyTerminalQuantumUses(Value value, - std::optional tensorSlot, - func::FuncOp func, - ArrayRef outputs, - DenseSet& visited) { - if (!visited.insert(value).second) { - return true; - } - return llvm::all_of(value.getUses(), [&](OpOperand& use) { - Operation* owner = use.getOwner(); - if (isa(owner)) { - return true; - } - if (auto measure = dyn_cast(owner)) { - return !tensorSlot && measure->getParentOfType() == func && - hasOutputOnlyMeasurementResult(measure, outputs) && - hasOnlyTerminalQuantumUses(measure.getQubitOut(), std::nullopt, - func, outputs, visited); - } - if (auto fromElements = dyn_cast(owner)) { - return !tensorSlot && hasOnlyTerminalQuantumUses(fromElements.getResult(), - use.getOperandNumber(), - func, outputs, visited); - } - if (auto insert = dyn_cast(owner)) { - const auto index = getConstantTensorIndex(insert.getIndex()); - if (!index) { - return false; - } - if (use.get() == insert.getScalar()) { - return !tensorSlot && - hasOnlyTerminalQuantumUses(insert.getResult(), index, func, - outputs, visited); - } - return tensorSlot && *tensorSlot != *index && - hasOnlyTerminalQuantumUses(insert.getResult(), tensorSlot, func, - outputs, visited); - } - if (auto extract = dyn_cast(owner)) { - const auto index = getConstantTensorIndex(extract.getIndex()); - if (!tensorSlot || !index) { - return false; - } - if (*tensorSlot == *index) { - return hasOnlyTerminalQuantumUses(extract.getResult(), std::nullopt, - func, outputs, visited); - } - return hasOnlyTerminalQuantumUses(extract.getOutTensor(), tensorSlot, - func, outputs, visited); - } - return false; - }); -} - -static bool isDeferrableMeasurement(MeasureOp measure, func::FuncOp func, +static bool isDeferrableMeasurement(MeasureOp measure, Block* entry, ArrayRef outputs) { - if (!hasOutputOnlyMeasurementResult(measure, outputs)) { + if (measure->getBlock() != entry || + !hasOutputOnlyMeasurementResult(measure, outputs)) { return false; } - DenseSet visited; - return hasOnlyTerminalQuantumUses(measure.getQubitOut(), std::nullopt, func, - outputs, visited); + return llvm::all_of(measure.getQubitOut().getUses(), [&](OpOperand& use) { + Operation* owner = use.getOwner(); + return isa(owner) || + (owner == entry->getTerminator() && isa(owner)); + }); } -static void analyzeSampling(func::FuncOp func, ArrayRef outputs, +static void analyzeSampling(func::FuncOp func, Block* sampledEntry, + ArrayRef outputs, DenseSet& active, SamplingPlan& plan) { Operation* funcOp = func.getOperation(); if (!active.insert(funcOp).second) { @@ -1973,7 +1919,7 @@ static void analyzeSampling(func::FuncOp func, ArrayRef outputs, if (isa(op)) { plan.dynamic = true; } else if (auto measure = dyn_cast(op)) { - if (isDeferrableMeasurement(measure, func, outputs)) { + if (isDeferrableMeasurement(measure, sampledEntry, outputs)) { plan.deferredMeasurements.insert(op); } else { plan.dynamic = true; @@ -1985,7 +1931,7 @@ static void analyzeSampling(func::FuncOp func, ArrayRef outputs, !callee.getBody().hasOneBlock()) { plan.dynamic = true; } else { - analyzeSampling(callee, outputs, active, plan); + analyzeSampling(callee, sampledEntry, outputs, active, plan); } } }); @@ -2011,7 +1957,7 @@ static FailureOr getSamplingPlan(func::FuncOp func) { } DenseSet active; - analyzeSampling(func, plan.outputs, active, plan); + analyzeSampling(func, &entry, plan.outputs, active, plan); return plan; } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 5cf49a7d10..eb4c4fb4c5 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -1881,6 +1881,39 @@ TEST_F(QCODDFunctionalityTest, ScfForSharesExecutionBudget) { EXPECT_TRUE(dd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, ExecutionBudgetIncludesBranchesAndCalls) { + for (const StringRef source : { + R"mlir(module { + func.func @main() { + %true = arith.constant true + %zero = arith.constant 0 : index + %limit = arith.constant 10000 : index + %one = arith.constant 1 : index + scf.for %i = %zero to %limit step %one { + scf.if %true { + } + } + return + } + })mlir", + R"mlir(module { + func.func @noop() { + return + } + func.func @main() { + %zero = arith.constant 0 : index + %limit = arith.constant 10000 : index + %one = arith.constant 1 : index + scf.for %i = %zero to %limit step %one { + func.call @noop() : () -> () + } + return + } + })mlir"}) { + expectMlirSimulationFails(0, source); + } +} + TEST_F(QCODDFunctionalityTest, SimulateRicherClassicalArithmetic) { auto mod = buildModule([](QCOProgramBuilder& b) { auto q = b.staticQubit(0); @@ -1986,39 +2019,6 @@ TEST_F(QCODDFunctionalityTest, EXPECT_TRUE(dd->getRootSet().empty()); } -TEST_F(QCODDFunctionalityTest, DefersTensorMeasurementDespiteLaterUnrelatedOp) { - auto mod = buildModule([](QCOProgramBuilder& b) { - auto reg = - b.allocClassicalBitRegister(1, {}, cbit::Initialization::Undefined); - auto q0 = b.h(b.allocQubit()); - auto q1 = b.allocQubit(); - std::tie(q0, std::ignore) = b.measure(q0, reg, 0); - auto tensor = b.qtensorFromElements({q0, q1}); - std::tie(tensor, q0) = b.qtensorExtract(tensor, 0); - tensor = b.qtensorInsert(q0, tensor, 0); - std::tie(tensor, q1) = b.qtensorExtract(tensor, 1); - q1 = b.x(q1); - tensor = b.qtensorInsert(q1, tensor, 1); - b.qtensorDealloc(tensor); - return reg; - }); - ASSERT_TRUE(mod); - - std::mt19937_64 rng(11); - auto singleDD = std::make_unique(2); - ASSERT_TRUE(succeeded(sample(mainFunc(*mod), *singleDD, 1, rng))); - const auto singleEvolutionLookups = - singleDD->matrixVectorMultiplication.getStats().lookups; - - auto dd = std::make_unique(2); - const auto histogram = sample(mainFunc(*mod), *dd, 64, rng); - ASSERT_TRUE(succeeded(histogram)); - EXPECT_EQ(histogram->at("0") + histogram->at("1"), 64U); - EXPECT_EQ(dd->matrixVectorMultiplication.getStats().lookups, - singleEvolutionLookups); - EXPECT_TRUE(dd->getRootSet().empty()); -} - TEST_F(QCODDFunctionalityTest, SampleDefersAllocatedQubitMeasurement) { auto mod = buildModule([](QCOProgramBuilder& b) { auto reg = From f4d08023da3141e8e94d11927a9e0b70e1e31147 Mon Sep 17 00:00:00 2001 From: Lukas Burgholzer Date: Mon, 31 Aug 2026 19:43:52 +0000 Subject: [PATCH 09/12] =?UTF-8?q?=E2=9C=85=20Cover=20optimized=20QCO=20whi?= =?UTF-8?q?le-loop=20sampling?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add an OpenQASM compiler-to-sampler regression for a measurement-controlled while loop and condense the QCO DD changelog entry. Assisted-by: GPT-5.6 via Codex --- CHANGELOG.md | 9 +++------ test/python/test_qco_dd.py | 19 +++++++++++++++++++ 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 55b7ec71d7..b0e4402d1f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,12 +32,9 @@ releases may include breaking changes. instance specifications, analytic references, deterministic manifests, and C++, Python, and command-line interfaces ([#2135]) ([**@denialhaag**], [**@burgholzer**]) -- ✨ Add decision diagram-based construction, simulation, and sampling for QCO - programs, including mid-circuit `measure` / `reset`, concrete QCO and SCF - control flow, non-recursive calls, bound parameters, classical integer and - `f64` SSA, CBit registers, one-dimensional memrefs, dynamic quantum allocation - and qtensors, dense multi-wire embedding, output-aware multi-shot sampling, - and Python bindings ([#1915], [#1973], [#2077], [#2078]) +- ✨ Add decision diagram-based construction, simulation, and output-aware + sampling for QCO programs, with C++ and Python APIs and support for classical + control flow and dynamic quantum data ([#1915], [#1973], [#2077], [#2078]) ([**@simon1hofmann**], [**@burgholzer**]) - ✨ Add immutable MLIR compiler targets, QDMI device integration, and target compilation through C++, Python, and `mqt-cc` ([#1687], [#1993], [#1999], diff --git a/test/python/test_qco_dd.py b/test/python/test_qco_dd.py index 2aa732115e..31714d8d9a 100644 --- a/test/python/test_qco_dd.py +++ b/test/python/test_qco_dd.py @@ -175,3 +175,22 @@ def test_compiler_to_sampler_outputs(source: str, num_qubits: int, expected: set assert set(counts) == expected assert sum(counts.values()) == shots + + +def test_compiler_to_sampler_while_reset() -> None: + """An imported measurement-controlled while loop terminates at zero.""" + source = """ +OPENQASM 3.0; +include "stdgates.inc"; +qubit q; +h q; +bit repeat = measure q; +while (repeat) { h q; repeat = measure q; } +output bit out; +out = measure q; +""" + program = compile_program(source, output=OutputFormat.QCO_OPTIMIZED) + package = DDPackage(1) + shots = 256 + + assert program.sample(package, shots=shots, seed=17) == {"0": shots} From 0060dac77faeb180707775aa2bf45cea441436d2 Mon Sep 17 00:00:00 2001 From: Lukas Burgholzer Date: Mon, 31 Aug 2026 20:51:38 +0000 Subject: [PATCH 10/12] =?UTF-8?q?=E2=9C=85=20Cover=20remaining=20QCO=20DD?= =?UTF-8?q?=20arithmetic=20paths?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Exercise the remaining producer-backed integer and floating-point handlers through the existing observable memref call path. Assisted-by: GPT-5.6 via Codex --- .../Dialect/QCO/Utils/test_dd_functionality.cpp | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index eb4c4fb4c5..10840b8479 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -2553,9 +2553,12 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { %two = arith.constant 2 : i16 %fourteen = arith.muli %seven, %two : i16 %quotient = arith.divsi %fourteen, %two : i16 + %unsigned_quotient = arith.divui %fourteen, %two : i16 %remainder = arith.remui %quotient, %two : i16 + %signed_remainder = arith.remsi %unsigned_quotient, %two : i16 %shifted = arith.shli %remainder, %two : i16 %restored = arith.shrui %shifted, %two : i16 + %signed_restored = arith.shrsi %shifted, %two : i16 %wide = arith.extui %restored : i16 to i32 %narrow = arith.trunci %wide : i32 to i16 %as_float = arith.sitofp %narrow : i16 to f64 @@ -2566,15 +2569,23 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { %expected = arith.constant 7 : i16 %integer_ok = arith.cmpi eq, %stored, %expected : i16 %casts_ok = arith.cmpi eq, %back, %remainder : i16 + %signed_remainder_ok = arith.cmpi eq, %signed_remainder, %remainder : i16 %one_float = arith.constant 1.0 : f64 + %negative_one_float = arith.negf %one_float : f64 %two_float = arith.addf %one_float, %one_float : f64 + %three_float = arith.subf %two_float, %negative_one_float : f64 + %six_float = arith.mulf %three_float, %two_float : f64 %four_float = arith.addf %two_float, %two_float : f64 %half = arith.divf %four_float, %two_float : f64 %float_remainder = arith.remf %half, %one_float : f64 %zero_float = arith.constant 0.0 : f64 - %float_ok = arith.cmpf oeq, %float_remainder, %zero_float : f64 + %six = arith.constant 6.0 : f64 + %remainder_ok = arith.cmpf oeq, %float_remainder, %zero_float : f64 + %product_ok = arith.cmpf oeq, %six_float, %six : f64 %integer_and_casts = arith.andi %integer_ok, %casts_ok : i1 - %condition = arith.andi %integer_and_casts, %float_ok : i1 + %all_integers_ok = arith.andi %integer_and_casts, %signed_remainder_ok : i1 + %floats_ok = arith.andi %remainder_ok, %product_ok : i1 + %condition = arith.andi %all_integers_ok, %floats_ok : i1 %q = qco.static 0 : !qco.qubit %q1 = qco.if %condition args(%qin = %q) -> (!qco.qubit) { %out = qco.x %qin : !qco.qubit -> !qco.qubit From e3dead8a9920354fa07de94bff7f97888564eb6b Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Tue, 1 Sep 2026 09:00:02 +0200 Subject: [PATCH 11/12] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Simplify=20QCO=20DD?= =?UTF-8?q?=20classical=20execution?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reuse MLIR folding for behavior-equivalent arithmetic while retaining explicit handling where fold semantics differ. Remove verifier-redundant guards and consolidate regression coverage. Assisted-by: GPT-5.6 via Codex Signed-off-by: Simon Hofmann --- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 328 ++++++------------ .../QCO/Utils/test_dd_functionality.cpp | 127 ++++--- test/python/test_qco_dd.py | 35 +- 3 files changed, 206 insertions(+), 284 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 5687bddfba..0494798289 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -711,9 +711,6 @@ static LogicalResult applyMemRefAlloc(memref::AllocOp alloc, } int64_t size = type.getDimSize(0); if (type.isDynamicDim(0)) { - if (alloc.getDynamicSizes().size() != 1) { - return alloc.emitError() << "dynamic 1-D memref requires one size"; - } auto dynamicSize = lookupIndex(alloc.getDynamicSizes()[0], classical, alloc); if (failed(dynamicSize)) { @@ -760,30 +757,6 @@ static LogicalResult applyMemRefLoad(memref::LoadOp load, return success(); } -template -static LogicalResult applyBinaryInteger(OpTy op, ClassicalEnv& classical, - Combine combine) { - auto lhs = lookupInteger(op.getLhs(), classical, op); - auto rhs = lookupInteger(op.getRhs(), classical, op); - if (failed(lhs) || failed(rhs)) { - return failure(); - } - return bindInteger(op.getResult(), combine(*lhs, *rhs), classical); -} - -template -static LogicalResult applyBinaryFloat(OpTy op, ClassicalEnv& classical, - Combine combine) { - auto lhs = lookupFloat(op.getLhs(), classical, op); - auto rhs = lookupFloat(op.getRhs(), classical, op); - if (failed(lhs) || failed(rhs)) { - return failure(); - } - classical.values[op.getResult()] = - FloatAttr::get(op.getResult().getType(), combine(*lhs, *rhs)); - return success(); -} - template static LogicalResult applyDivision(OpTy op, ClassicalEnv& classical, Combine combine) { @@ -801,20 +774,6 @@ static LogicalResult applyDivision(OpTy op, ClassicalEnv& classical, return bindInteger(op.getResult(), combine(*lhs, *rhs), classical); } -template -static LogicalResult applyShift(OpTy op, ClassicalEnv& classical, Shift shift) { - auto lhs = lookupInteger(op.getLhs(), classical, op); - auto rhs = lookupInteger(op.getRhs(), classical, op); - if (failed(lhs) || failed(rhs)) { - return failure(); - } - if (rhs->uge(lhs->getBitWidth())) { - return op.emitError() << "shift amount out of range for QCO DD simulation"; - } - return bindInteger(op.getResult(), shift(*lhs, rhs->getZExtValue()), - classical); -} - static LogicalResult applyIntegerCast(Value in, Value out, Operation* op, ClassicalEnv& classical, bool isSigned) { auto value = lookupInteger(in, classical, op); @@ -832,6 +791,98 @@ static LogicalResult applyIntegerCast(Value in, Value out, Operation* op, return bindInteger(out, *value, classical); } +static LogicalResult foldClassicalOp(Operation& op, ClassicalEnv& classical) { + if (llvm::any_of(op.getOperandTypes(), + [](Type type) { return !isSupportedClassicalType(type); }) || + llvm::any_of(op.getResultTypes(), + [](Type type) { return !isSupportedClassicalType(type); })) { + return op.emitError() + << "QCO DD simulation only supports integer, index, and f64 values"; + } + + const auto lookupOperands = + [&](Operation& candidate) -> FailureOr> { + SmallVector operands; + operands.reserve(candidate.getNumOperands()); + for (Value operand : candidate.getOperands()) { + auto attr = lookupAttribute(operand, classical, &op); + if (failed(attr)) { + return failure(); + } + const auto typed = dyn_cast(*attr); + if (!typed || typed.getType() != operand.getType()) { + return op.emitError() + << "classical SSA value has a mismatched runtime attribute"; + } + operands.push_back(*attr); + } + return operands; + }; + + auto operands = lookupOperands(op); + if (failed(operands)) { + return failure(); + } + + if (isa(op)) { + const auto lhs = dyn_cast((*operands)[0]); + const auto rhs = dyn_cast((*operands)[1]); + if (!lhs || !rhs) { + return op.emitError() << "expected integer shift operands"; + } + if (rhs.getValue().uge(lhs.getValue().getBitWidth())) { + return op.emitError() + << "shift amount out of range for QCO DD simulation"; + } + } + + // Fold a clone because some arithmetic folders canonicalize in place. + Operation* clone = op.clone(); + const auto destroyClone = llvm::make_scope_exit([&] { clone->destroy(); }); + for (unsigned attempt = 0; attempt < 2; ++attempt) { + SmallVector results; + if (failed(clone->fold(*operands, results))) { + return op.emitError() + << "could not evaluate classical op during QCO DD simulation"; + } + + Attribute result; + if (results.size() == 1) { + result = dyn_cast_if_present(results.front()); + if (!result) { + const auto value = dyn_cast_if_present(results.front()); + if (value && value != clone->getResult(0)) { + auto attr = lookupAttribute(value, classical, &op); + if (failed(attr)) { + return failure(); + } + result = *attr; + } + } + } + if (result) { + const auto typed = dyn_cast(result); + if (!typed || typed.getType() != op.getResult(0).getType()) { + return op.emitError() + << "folded classical value has a mismatched result type"; + } + classical.values[op.getResult(0)] = result; + return success(); + } + if (attempt == 0 && (results.empty() || results.size() == 1)) { + operands = lookupOperands(*clone); + if (failed(operands)) { + return failure(); + } + continue; + } + return op.emitError() + << "could not evaluate classical op during QCO DD simulation"; + } + return op.emitError() + << "could not evaluate classical op during QCO DD simulation"; +} + static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { const auto isUnsupportedFloat = [](Type type) { return isa(type) && !type.isF64(); @@ -842,102 +893,42 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { << "QCO DD simulation only supports f64 classical values"; } return TypeSwitch(&op) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs & rhs; - }); - }) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs | rhs; - }); - }) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs ^ rhs; - }); - }) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs + rhs; - }); - }) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs - rhs; - }); - }) - .Case([&](auto value) { - return applyBinaryInteger( - value, classical, - [](const llvm::APInt& lhs, const llvm::APInt& rhs) { - return lhs * rhs; - }); - }) - .Case([&](auto value) { + .Case( + [&](Operation* foldable) { + return foldClassicalOp(*foldable, classical); + }) + .Case([&](arith::DivUIOp value) { return applyDivision( value, classical, [](const llvm::APInt& lhs, const llvm::APInt& rhs) { return lhs.udiv(rhs); }); }) - .Case([&](auto value) { + .Case([&](arith::DivSIOp value) { return applyDivision( value, classical, [](const llvm::APInt& lhs, const llvm::APInt& rhs) { return lhs.sdiv(rhs); }); }) - .Case([&](auto value) { + .Case([&](arith::RemUIOp value) { return applyDivision( value, classical, [](const llvm::APInt& lhs, const llvm::APInt& rhs) { return lhs.urem(rhs); }); }) - .Case([&](auto value) { + .Case([&](arith::RemSIOp value) { return applyDivision( value, classical, [](const llvm::APInt& lhs, const llvm::APInt& rhs) { return lhs.srem(rhs); }); }) - .Case([&](auto value) { - return applyShift( - value, classical, - [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.shl(rhs); }); - }) - .Case([&](auto value) { - return applyShift( - value, classical, - [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.lshr(rhs); }); - }) - .Case([&](auto value) { - return applyShift( - value, classical, - [](const llvm::APInt& lhs, uint64_t rhs) { return lhs.ashr(rhs); }); - }) - .Case([&](arith::CmpIOp cmp) -> LogicalResult { - auto lhs = lookupInteger(cmp.getLhs(), classical, cmp); - auto rhs = lookupInteger(cmp.getRhs(), classical, cmp); - if (failed(lhs) || failed(rhs)) { - return failure(); - } - classical.values[cmp.getResult()] = BoolAttr::get( - cmp.getContext(), - arith::applyCmpPredicate(cmp.getPredicate(), *lhs, *rhs)); - return success(); - }) .Case([&](arith::SelectOp select) -> LogicalResult { auto condition = lookupBool(select.getCondition(), classical, select); if (failed(condition)) { @@ -967,60 +958,6 @@ static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, false); }) - .Case([&](auto value) { - return applyBinaryFloat( - value, classical, [](double lhs, double rhs) { return lhs + rhs; }); - }) - .Case([&](auto value) { - return applyBinaryFloat( - value, classical, [](double lhs, double rhs) { return lhs - rhs; }); - }) - .Case([&](auto value) { - return applyBinaryFloat( - value, classical, [](double lhs, double rhs) { return lhs * rhs; }); - }) - .Case([&](auto value) { - return applyBinaryFloat( - value, classical, [](double lhs, double rhs) { return lhs / rhs; }); - }) - .Case([&](auto value) { - return applyBinaryFloat(value, classical, [](double lhs, double rhs) { - return std::fmod(lhs, rhs); - }); - }) - .Case([&](arith::NegFOp neg) -> LogicalResult { - auto value = lookupFloat(neg.getOperand(), classical, neg); - if (failed(value)) { - return failure(); - } - classical.values[neg.getResult()] = - FloatAttr::get(neg.getType(), -*value); - return success(); - }) - .Case([&](arith::CmpFOp cmp) -> LogicalResult { - auto lhs = lookupFloat(cmp.getLhs(), classical, cmp); - auto rhs = lookupFloat(cmp.getRhs(), classical, cmp); - if (failed(lhs) || failed(rhs)) { - return failure(); - } - classical.values[cmp.getResult()] = BoolAttr::get( - cmp.getContext(), - arith::applyCmpPredicate(cmp.getPredicate(), llvm::APFloat(*lhs), - llvm::APFloat(*rhs))); - return success(); - }) - .Case( - [&](Operation* castOp) -> LogicalResult { - auto value = - lookupInteger(castOp->getOperand(0), classical, castOp); - if (failed(value)) { - return failure(); - } - classical.values[castOp->getResult(0)] = FloatAttr::get( - castOp->getResult(0).getType(), - value->roundToDouble(isa(castOp))); - return success(); - }) .Case( [&](Operation* castOp) -> LogicalResult { auto value = lookupFloat(castOp->getOperand(0), classical, castOp); @@ -1083,17 +1020,6 @@ static FailureOr resolveLoop(scf::ForOp forOp, static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, WalkState& walk, Operation* op); -static LogicalResult bindLinearArgs(ValueRange operands, Block& block, - WalkState& walk, Operation* op) { - for (Value arg : block.getArguments()) { - if (!isa(arg.getType()) && !isQTensorType(arg.getType())) { - return op->emitError() - << "unsupported linear region argument for QCO DD simulation"; - } - } - return bindValuePairs(operands, block.getArguments(), walk, op); -} - static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, WalkState& walk, Operation* op) { SmallVector values; @@ -1160,10 +1086,6 @@ static LogicalResult bindYieldResults(YieldOp yield, ValueRange linearResults, WalkState& walk) { const size_t numClassical = classicalResults.size(); - if (yield.getNumOperands() != numClassical + linearResults.size()) { - return yield.emitError() - << "yield operand count does not match operation results"; - } if (failed(bindValuePairs(yield.getOperands().take_front(numClassical), classicalResults, walk, yield))) { return failure(); @@ -1190,7 +1112,8 @@ static LogicalResult applyRegionBranch(ValueRange linearOperands, Block& block, ValueRange classicalResults, ValueRange linearResults, WalkState& walk, StateDD& state, Operation* parent) { - if (failed(bindLinearArgs(linearOperands, block, walk, parent))) { + if (failed( + bindValuePairs(linearOperands, block.getArguments(), walk, parent))) { return failure(); } if (failed(consumeExecutionStep(walk, parent))) { @@ -1229,9 +1152,6 @@ static LogicalResult applyScfRegion(Region& region, ValueRange results, static FailureOr allocateZeroQubits(size_t count, WalkState& walk, dd::VectorDD& state, Operation* op) { - if (count == 0) { - return op->emitError() << "quantum allocation size must be positive"; - } if (walk.qubits->numQubits > walk.dd->qubits() || count > walk.dd->qubits() - walk.qubits->numQubits) { return op->emitError() << "DD package has " << walk.dd->qubits() @@ -1395,18 +1315,6 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return storeRegister(store, *walk.classical); }) .template Case([](auto) { return success(); }) - .template Case< - arith::AndIOp, arith::OrIOp, arith::XOrIOp, arith::AddIOp, - arith::SubIOp, arith::MulIOp, arith::DivUIOp, arith::DivSIOp, - arith::RemUIOp, arith::RemSIOp, arith::ShLIOp, arith::ShRUIOp, - arith::ShRSIOp, arith::CmpIOp, arith::SelectOp, arith::ExtUIOp, - arith::ExtSIOp, arith::IndexCastUIOp, arith::IndexCastOp, - arith::TruncIOp, arith::AddFOp, arith::SubFOp, arith::MulFOp, - arith::DivFOp, arith::RemFOp, arith::NegFOp, arith::CmpFOp, - arith::SIToFPOp, arith::UIToFPOp, arith::FPToSIOp, arith::FPToUIOp>( - [&](Operation* classicalOp) { - return applyClassicalOp(*classicalOp, *walk.classical); - }) .template Case([&](func::ReturnOp returnOp) { return validateReturn(returnOp, *walk.qubits, *walk.tensors); }) @@ -1471,9 +1379,6 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return failure(); } Block* block = *condition ? ifOp.thenBlock() : ifOp.elseBlock(); - if (block == nullptr) { - return ifOp.emitError() << "selected qco.if region is empty"; - } return applyRegionBranch(ifOp.getQubits(), *block, ifOp.getClassicalResults(), ifOp.getLinearResults(), walk, state, ifOp); @@ -1492,10 +1397,6 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { break; } } - if (block == nullptr) { - return switchOp.emitError() - << "selected qco.index_switch region is empty"; - } return applyRegionBranch( switchOp.getTargets(), *block, switchOp.getClassicalResults(), switchOp.getLinearResults(), walk, state, switchOp); @@ -1572,11 +1473,6 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { return bindValuePairs(carried, forOp.getResults(), walk, forOp); }) .template Case([&](scf::WhileOp whileOp) -> LogicalResult { - if (!whileOp.getBefore().hasOneBlock() || - !whileOp.getAfter().hasOneBlock()) { - return whileOp.emitError() - << "scf.while regions must contain one block"; - } Block& before = whileOp.getBefore().front(); Block& after = whileOp.getAfter().front(); SmallVector carried(whileOp.getInits().begin(), @@ -1587,11 +1483,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { failed(walkBlock(before, walk, state))) { return failure(); } - auto condition = dyn_cast(before.getTerminator()); - if (!condition) { - return whileOp.emitError() - << "scf.while before region missing scf.condition"; - } + auto condition = whileOp.getConditionOp(); auto value = lookupBool(condition.getCondition(), *walk.classical, whileOp); if (failed(value)) { @@ -1609,11 +1501,7 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { failed(walkBlock(after, walk, state))) { return failure(); } - auto yield = dyn_cast(after.getTerminator()); - if (!yield) { - return whileOp.emitError() - << "scf.while after region missing scf.yield"; - } + auto yield = whileOp.getYieldOp(); carried.assign(yield.getOperands().begin(), yield.getOperands().end()); } diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index 10840b8479..df16c7497a 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -25,6 +25,7 @@ #include #include +#include #include #include #include @@ -956,12 +957,15 @@ TEST_F(QCODDFunctionalityTest, AcceptsLargestValidShift) { } TEST_F(QCODDFunctionalityTest, RejectsOutOfRangeShift) { - for (const int64_t amount : {-1, 64}) { - auto mod = buildModule([amount](QCOProgramBuilder& b) { + for (const auto [lhs, amount] : + {std::pair{1, -1}, {1, 64}, {0, 64}, {64, 64}}) { + auto mod = buildModule([lhs, amount](QCOProgramBuilder& b) { auto q = b.staticQubit(0); - auto one = arith::ConstantIndexOp::create(b, 1).getResult(); - auto bad = arith::ConstantIndexOp::create(b, amount).getResult(); - auto shifted = arith::ShLIOp::create(b, one, bad).getResult(); + auto value = arith::ConstantIndexOp::create(b, lhs).getResult(); + auto bad = lhs == amount + ? value + : arith::ConstantIndexOp::create(b, amount).getResult(); + auto shifted = arith::ShLIOp::create(b, value, bad).getResult(); q = b.qcoIndexSwitch(shifted, q, ArrayRef{0}, SmallVector>{ [&](Value arg) { return arg; }}, @@ -1202,12 +1206,7 @@ TEST_F(QCODDFunctionalityTest, RejectsUnsupportedOrUnboundClassicalOperations) { return } })mlir"}) { - auto mod = parseSourceString(source, context.get()); - 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))); + expectMlirSimulationFails(1, source); } } TEST_F(QCODDFunctionalityTest, RejectsUnmappedClassicalControl) { @@ -1358,12 +1357,7 @@ TEST_F(QCODDFunctionalityTest, RejectsUnboundClassicalRegionResults) { return } })mlir"}) { - auto mod = parseSourceString(source, context.get()); - 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))); + expectMlirSimulationFails(1, source); } } TEST_F(QCODDFunctionalityTest, Rejects) { @@ -2031,7 +2025,6 @@ TEST_F(QCODDFunctionalityTest, SampleDefersAllocatedQubitMeasurement) { ASSERT_TRUE(mod); auto dd = std::make_unique(1); - std::mt19937_64 rng(11); const auto histogram = sample(mainFunc(*mod), *dd, 8, rng); ASSERT_TRUE(succeeded(histogram)); EXPECT_EQ(*histogram, (std::map{{"1", 8}})); @@ -2255,7 +2248,6 @@ TEST_F(QCODDFunctionalityTest, SymbolicParametersUseBindings) { dd->decRef(*actual); dd->decRef(*expected); - std::mt19937_64 rng(5); const auto histogram = sample(func, *dd, 8, rng, bindings); ASSERT_TRUE(succeeded(histogram)); EXPECT_EQ(*histogram, (std::map{{"1", 8}})); @@ -2419,7 +2411,6 @@ TEST_F(QCODDFunctionalityTest, DynamicAllocationsAndQTensorBookkeeping) { ASSERT_TRUE(mod); auto dd = std::make_unique(2); - std::mt19937_64 rng(3); const auto histogram = sample(mainFunc(*mod), *dd, 8, rng); ASSERT_TRUE(succeeded(histogram)); EXPECT_EQ(*histogram, (std::map{{"11", 8}})); @@ -2469,7 +2460,6 @@ TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { IntegerAttr::get(IndexType::get(context.get()), 2); auto dd = std::make_unique(2); - std::mt19937_64 rng(7); const auto histogram = sample(func, *dd, 4, rng, bindings); ASSERT_TRUE(succeeded(histogram)); EXPECT_EQ(*histogram, (std::map{{"10", 4}})); @@ -2484,18 +2474,13 @@ TEST_F(QCODDFunctionalityTest, DynamicQTensorArgumentUsesBoundExtent) { } TEST_F(QCODDFunctionalityTest, RejectsQTensorBeyondQubitRange) { - auto mod = parseSourceString(R"mlir( + expectMlirFails(1, R"mlir( module { func.func @main(%qubits: tensor<65537x!qco.qubit>) { return } } - )mlir", - context.get()); - ASSERT_TRUE(mod); - - auto dd = std::make_unique(1); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*mod), *dd))); + )mlir"); } TEST_F(QCODDFunctionalityTest, QTensorFlowsThroughLoopAndCall) { @@ -2530,7 +2515,6 @@ TEST_F(QCODDFunctionalityTest, QTensorFlowsThroughLoopAndCall) { ASSERT_TRUE(mod); auto dd = std::make_unique(1); - std::mt19937_64 rng(13); const auto histogram = sample(mainFunc(*mod), *dd, 4, rng); ASSERT_TRUE(succeeded(histogram)); EXPECT_EQ(*histogram, (std::map{{"1", 4}})); @@ -2571,6 +2555,8 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { %casts_ok = arith.cmpi eq, %back, %remainder : i16 %signed_remainder_ok = arith.cmpi eq, %signed_remainder, %remainder : i16 %one_float = arith.constant 1.0 : f64 + %unsigned_float = arith.uitofp %narrow : i16 to f64 + %unsigned_float_ok = arith.cmpf oeq, %unsigned_float, %one_float : f64 %negative_one_float = arith.negf %one_float : f64 %two_float = arith.addf %one_float, %one_float : f64 %three_float = arith.subf %two_float, %negative_one_float : f64 @@ -2584,7 +2570,8 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { %product_ok = arith.cmpf oeq, %six_float, %six : f64 %integer_and_casts = arith.andi %integer_ok, %casts_ok : i1 %all_integers_ok = arith.andi %integer_and_casts, %signed_remainder_ok : i1 - %floats_ok = arith.andi %remainder_ok, %product_ok : i1 + %float_results_ok = arith.andi %remainder_ok, %product_ok : i1 + %floats_ok = arith.andi %float_results_ok, %unsigned_float_ok : i1 %condition = arith.andi %all_integers_ok, %floats_ok : i1 %q = qco.static 0 : !qco.qubit %q1 = qco.if %condition args(%qin = %q) -> (!qco.qubit) { @@ -2607,7 +2594,7 @@ TEST_F(QCODDFunctionalityTest, WiderMemRefCallsShareStorage) { TEST_F(QCODDFunctionalityTest, BindingsDriveObservableClassicalPath) { auto mod = parseSourceString(R"mlir( module { - func.func @main(%idx: index, %word: i16, %flag: i1) { + func.func @main(%idx: index, %word: i16, %flag: i1, %unbound: i1) { %one = arith.constant 1 : i16 %expected = arith.constant 4 : i16 %sum = arith.addi %word, %one : i16 @@ -2615,7 +2602,7 @@ TEST_F(QCODDFunctionalityTest, BindingsDriveObservableClassicalPath) { %as_index = arith.index_cast %word : i16 to index %index_ok = arith.cmpi eq, %as_index, %idx : index %both = arith.andi %integer_ok, %index_ok : i1 - %condition = arith.select %flag, %both, %flag : i1 + %condition = arith.select %flag, %both, %unbound : i1 %q = qco.static 0 : !qco.qubit %out = qco.if %condition args(%arg = %q) -> (!qco.qubit) { %flipped = qco.x %arg : !qco.qubit -> !qco.qubit @@ -2650,6 +2637,44 @@ TEST_F(QCODDFunctionalityTest, BindingsDriveObservableClassicalPath) { EXPECT_TRUE(dd->getRootSet().empty()); } +TEST_F(QCODDFunctionalityTest, RepeatedSimulationPreservesFoldableIR) { + auto mod = parseSourceString(R"mlir( + module { + func.func @main() { + %zero = arith.constant 0 : i8 + %wide = arith.extui %zero : i8 to i16 + %narrow = arith.trunci %wide : i16 to i8 + %wide_again = arith.extsi %narrow : i8 to i16 + %equal = arith.cmpi eq, %wide, %wide_again : i16 + %q = qco.static 0 : !qco.qubit + %out = qco.if %equal args(%arg = %q) -> (!qco.qubit) { + %flipped = qco.x %arg : !qco.qubit -> !qco.qubit + qco.yield %flipped : !qco.qubit + } else args(%arg = %q) { + qco.yield %arg : !qco.qubit + } + qco.sink %out : !qco.qubit + return + } + } + )mlir", + context.get()); + ASSERT_TRUE(mod); + + const auto print = [](ModuleOp module) { + std::string result; + llvm::raw_string_ostream stream(result); + module.print(stream); + stream.flush(); + return result; + }; + const auto original = print(*mod); + expectSimulatesFromZero(mainFunc(*mod), true); + EXPECT_EQ(print(*mod), original); + expectSimulatesFromZero(mainFunc(*mod), true); + EXPECT_EQ(print(*mod), original); +} + TEST_F(QCODDFunctionalityTest, RejectsRepresentativeClassicalRuntimeErrors) { for (const StringRef source : { R"mlir(module { @@ -2685,8 +2710,28 @@ TEST_F(QCODDFunctionalityTest, RejectsRepresentativeClassicalRuntimeErrors) { R"mlir(module { func.func @main() { %zero = arith.constant 0 : i8 - %one = arith.constant 1 : i8 - %invalid = arith.divui %one, %zero : i8 + %invalid = arith.divui %zero, %zero : i8 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %zero = arith.constant 0 : i8 + %invalid = arith.divsi %zero, %zero : i8 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %zero = arith.constant 0 : i8 + %invalid = arith.remui %zero, %zero : i8 + return + } + })mlir", + R"mlir(module { + func.func @main() { + %zero = arith.constant 0 : i8 + %invalid = arith.remsi %zero, %zero : i8 return } })mlir", @@ -2739,7 +2784,7 @@ TEST_F(QCODDFunctionalityTest, BuildFunctionalityRestrictsRuntimeAllocations) { } )mlir", context.get()); - auto nested = parseSourceString(R"mlir( + expectMlirFails(1, R"mlir( module { func.func @main() { %true = arith.constant true @@ -2750,9 +2795,8 @@ TEST_F(QCODDFunctionalityTest, BuildFunctionalityRestrictsRuntimeAllocations) { return } } - )mlir", - context.get()); - auto tensor = parseSourceString(R"mlir( + )mlir"); + expectMlirFails(1, R"mlir( module { func.func @main() { %one = arith.constant 1 : index @@ -2761,18 +2805,13 @@ TEST_F(QCODDFunctionalityTest, BuildFunctionalityRestrictsRuntimeAllocations) { return } } - )mlir", - context.get()); + )mlir"); ASSERT_TRUE(topLevel); - ASSERT_TRUE(nested); - ASSERT_TRUE(tensor); auto dd = std::make_unique(1); const auto functionality = buildFunctionality(mainFunc(*topLevel), *dd); ASSERT_TRUE(succeeded(functionality)); dd->decRef(*functionality); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*nested), *dd))); - EXPECT_TRUE(failed(buildFunctionality(mainFunc(*tensor), *dd))); } } // namespace diff --git a/test/python/test_qco_dd.py b/test/python/test_qco_dd.py index 31714d8d9a..fecf2d79f1 100644 --- a/test/python/test_qco_dd.py +++ b/test/python/test_qco_dd.py @@ -162,8 +162,22 @@ def test_entry_func_required() -> None: 1, {"00", "01"}, ), + ( + """ +OPENQASM 3.0; +include "stdgates.inc"; +qubit q; +h q; +bit repeat = measure q; +while (repeat) { h q; repeat = measure q; } +output bit out; +out = measure q; +""", + 1, + {"0"}, + ), ], - ids=["terminal-bell", "adaptive-reset"], + ids=["terminal-bell", "adaptive-reset", "while-reset"], ) def test_compiler_to_sampler_outputs(source: str, num_qubits: int, expected: set[str]) -> None: """Compile optimized QCO and sample the declared CBit output.""" @@ -175,22 +189,3 @@ def test_compiler_to_sampler_outputs(source: str, num_qubits: int, expected: set assert set(counts) == expected assert sum(counts.values()) == shots - - -def test_compiler_to_sampler_while_reset() -> None: - """An imported measurement-controlled while loop terminates at zero.""" - source = """ -OPENQASM 3.0; -include "stdgates.inc"; -qubit q; -h q; -bit repeat = measure q; -while (repeat) { h q; repeat = measure q; } -output bit out; -out = measure q; -""" - program = compile_program(source, output=OutputFormat.QCO_OPTIMIZED) - package = DDPackage(1) - shots = 256 - - assert program.sample(package, shots=shots, seed=17) == {"0": shots} From 5e8065b6b48eb693f887637ecd36561415f66870 Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Tue, 1 Sep 2026 09:43:35 +0200 Subject: [PATCH 12/12] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20Trim=20unreachable?= =?UTF-8?q?=20QCO=20DD=20fold=20checks?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Trust verifier and runtime-map type invariants while preserving MLIR folding retries. Cover the valid multi-block scf.execute_region rejection. Assisted-by: GPT-5.6 via Codex Signed-off-by: Simon Hofmann --- .../lib/Dialect/QCO/Utils/DDFunctionality.cpp | 85 ++++++------------- .../QCO/Utils/test_dd_functionality.cpp | 15 ++++ 2 files changed, 40 insertions(+), 60 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp index 0494798289..d0433eb183 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -800,87 +800,52 @@ static LogicalResult foldClassicalOp(Operation& op, ClassicalEnv& classical) { << "QCO DD simulation only supports integer, index, and f64 values"; } - const auto lookupOperands = - [&](Operation& candidate) -> FailureOr> { + // Fold a clone because some arithmetic folders canonicalize in place. + Operation* clone = op.clone(); + const auto destroyClone = llvm::make_scope_exit([&] { clone->destroy(); }); + Attribute result; + for (unsigned attempt = 0; attempt < 2 && !result; ++attempt) { SmallVector operands; - operands.reserve(candidate.getNumOperands()); - for (Value operand : candidate.getOperands()) { + operands.reserve(clone->getNumOperands()); + for (Value operand : clone->getOperands()) { auto attr = lookupAttribute(operand, classical, &op); if (failed(attr)) { return failure(); } - const auto typed = dyn_cast(*attr); - if (!typed || typed.getType() != operand.getType()) { - return op.emitError() - << "classical SSA value has a mismatched runtime attribute"; - } operands.push_back(*attr); } - return operands; - }; - - auto operands = lookupOperands(op); - if (failed(operands)) { - return failure(); - } - if (isa(op)) { - const auto lhs = dyn_cast((*operands)[0]); - const auto rhs = dyn_cast((*operands)[1]); - if (!lhs || !rhs) { - return op.emitError() << "expected integer shift operands"; - } - if (rhs.getValue().uge(lhs.getValue().getBitWidth())) { - return op.emitError() - << "shift amount out of range for QCO DD simulation"; + if (isa(op)) { + const auto lhs = cast(operands[0]); + const auto rhs = cast(operands[1]); + if (rhs.getValue().uge(lhs.getValue().getBitWidth())) { + return op.emitError() + << "shift amount out of range for QCO DD simulation"; + } } - } - // Fold a clone because some arithmetic folders canonicalize in place. - Operation* clone = op.clone(); - const auto destroyClone = llvm::make_scope_exit([&] { clone->destroy(); }); - for (unsigned attempt = 0; attempt < 2; ++attempt) { SmallVector results; - if (failed(clone->fold(*operands, results))) { - return op.emitError() - << "could not evaluate classical op during QCO DD simulation"; + if (failed(clone->fold(operands, results))) { + break; } - - Attribute result; if (results.size() == 1) { result = dyn_cast_if_present(results.front()); if (!result) { - const auto value = dyn_cast_if_present(results.front()); - if (value && value != clone->getResult(0)) { - auto attr = lookupAttribute(value, classical, &op); - if (failed(attr)) { - return failure(); - } - result = *attr; + const auto value = cast(results.front()); + if (value != clone->getResult(0)) { + return classical.bindFrom(value, op.getResult(0), &op); } } + } else if (!results.empty()) { + break; } - if (result) { - const auto typed = dyn_cast(result); - if (!typed || typed.getType() != op.getResult(0).getType()) { - return op.emitError() - << "folded classical value has a mismatched result type"; - } - classical.values[op.getResult(0)] = result; - return success(); - } - if (attempt == 0 && (results.empty() || results.size() == 1)) { - operands = lookupOperands(*clone); - if (failed(operands)) { - return failure(); - } - continue; - } + } + if (!result) { return op.emitError() << "could not evaluate classical op during QCO DD simulation"; } - return op.emitError() - << "could not evaluate classical op during QCO DD simulation"; + classical.values[op.getResult(0)] = result; + return success(); } static LogicalResult applyClassicalOp(Operation& op, ClassicalEnv& classical) { diff --git a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp index df16c7497a..5210e33d5e 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -2395,6 +2395,21 @@ TEST_F(QCODDFunctionalityTest, StandardScfRegionsAndWhileCarryValues) { expectEqualToQc(mainFunc(*mod), qc); } +TEST_F(QCODDFunctionalityTest, RejectsMultiBlockScfExecuteRegion) { + expectMlirSimulationFails(0, R"mlir( + module { + func.func @main() { + scf.execute_region { + scf.yield + ^next: + scf.yield + } + return + } + } + )mlir"); +} + TEST_F(QCODDFunctionalityTest, DynamicAllocationsAndQTensorBookkeeping) { auto mod = buildModule([](QCOProgramBuilder& b) { auto q0 = b.x(b.allocQubit());