diff --git a/CHANGELOG.md b/CHANGELOG.md index de2bda2db8..b0e4402d1f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,11 +32,10 @@ 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 `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**]) +- ✨ 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], [#2049]) ([**@MatthiasReumann**], [**@simon1hofmann**], [**@burgholzer**]) @@ -921,6 +920,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..219996f787 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. + * + * 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 DDArgumentBindings = DenseMap; + /** * @brief Sequentially build a matrix DD for a static unitary QCO `func.func`. * @@ -31,30 +43,37 @@ 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 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); +FailureOr buildFunctionality( + func::FuncOp func, dd::Package& dd, + const DDArgumentBindings& argumentBindings = DDArgumentBindings()); /** * @brief Simulate a QCO `func.func` that may contain measurements, resets, and @@ -63,17 +82,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 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 * `qco::verifyLinearity`. @@ -83,11 +100,15 @@ 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 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); +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`. @@ -99,7 +120,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 +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 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); - +sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng, + 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 a9e52d540e..d0433eb183 100644 --- a/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp +++ b/mlir/lib/Dialect/QCO/Utils/DDFunctionality.cpp @@ -29,14 +29,19 @@ #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 @@ -59,11 +64,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 +115,57 @@ 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; + + void bind(Value value, TensorState slots) { + tensors[value] = std::move(slots); + } + + [[nodiscard]] TensorState lookup(Value value) const { + const auto it = tensors.find(value); + 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 { struct RegisterBit { std::optional value; std::optional deferredWire; }; using RegisterState = std::vector; + using MemRefState = SmallVector; - 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,13 +175,18 @@ 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; }; + +using RuntimeValue = std::variant, + std::shared_ptr>; struct LoopRange { llvm::APInt induction, step; size_t trips; @@ -145,10 +198,50 @@ 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 && + isa(tensorType.getElementType()); +} + +static FailureOr +lookupAttribute(Value value, const ClassicalEnv& classical, Operation* op) { + const auto it = classical.values.find(value); + if (it == classical.values.end()) { + return op->emitError() + << "classical SSA value is not mapped for QCO DD simulation"; + } + 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()) { + if (const auto floating = dyn_cast(it->second)) { + return floating.getValue().convertToDouble(); + } + } + 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 +278,17 @@ 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(); + } + 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)}; } @@ -250,20 +345,27 @@ 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(); + } + 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)); + 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 +445,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) { + 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 +485,111 @@ 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); +[[nodiscard]] static bool isSupportedClassicalType(Type type) { + return isa(type) || type.isF64(); } -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 recordConstant(arith::ConstantOp constant, + ClassicalEnv& classical) { + 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 FailureOr lookupBool(Value value, ClassicalEnv& classical, +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; + } + ++boundArguments; + const Attribute attr = binding->second; + const Type type = argument.getType(); + if (isQTensorType(type)) { + 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) { - auto result = lookupInteger(value, classical, op); - if (failed(result)) { + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer || !value.getType().isInteger(1)) { return failure(); } - return !result->isZero(); + return !integer.getValue().isZero(); } -static FailureOr lookupIndex(Value value, ClassicalEnv& classical, - Operation* op) { - auto result = lookupInteger(value, classical, op); - if (failed(result)) { +static FailureOr +lookupIndex(Value value, const ClassicalEnv& classical, Operation* op) { + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer || !isa(value.getType())) { return failure(); } - return result->getSExtValue(); + return integer.getValue().getSExtValue(); } -static LogicalResult applyUnsignedIndexCast(Value in, Value out, Operation* op, - ClassicalEnv& classical) { - auto value = lookupInteger(in, classical, op); - if (failed(value)) { +static FailureOr lookupFloat(Value value, const ClassicalEnv& classical, + Operation* op) { + auto attr = lookupAttribute(value, classical, op); + const auto floating = + succeeded(attr) ? dyn_cast(*attr) : FloatAttr{}; + if (!floating || !value.getType().isF64()) { return failure(); } - const auto integerType = dyn_cast(out.getType()); - const unsigned width = integerType ? integerType.getWidth() : 64U; - bindInteger(out, value->zextOrTrunc(width), classical); + return floating.getValue().convertToDouble(); +} + +static FailureOr +lookupInteger(Value value, const ClassicalEnv& classical, Operation* op) { + if (!isa(value.getType())) { + return op->emitError() << "expected an integer or index SSA value"; + } + auto attr = lookupAttribute(value, classical, op); + const auto integer = + succeeded(attr) ? dyn_cast(*attr) : IntegerAttr{}; + if (!integer) { + return failure(); + } + return integer.getValue(); +} + +static LogicalResult bindInteger(Value dest, const llvm::APInt& value, + ClassicalEnv& classical) { + 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(); } @@ -428,7 +608,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 +665,284 @@ 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); +} + +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 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)) { + 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)); 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); - if (failed(lhs) || failed(rhs)) { +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(); + } + if (!**slot) { + return load.emitError() + << "read from an uninitialized classical memref element"; + } + classical.values[load.getResult()] = **slot; + 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); +} + +static LogicalResult applyIntegerCast(Value in, Value out, Operation* op, + ClassicalEnv& classical, bool isSigned) { + auto value = lookupInteger(in, classical, op); + if (failed(value)) { 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 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 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"; + } + + // 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(clone->getNumOperands()); + for (Value operand : clone->getOperands()) { + auto attr = lookupAttribute(operand, classical, &op); + if (failed(attr)) { + return failure(); + } + operands.push_back(*attr); + } + + 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"; + } + } + + SmallVector results; + if (failed(clone->fold(operands, results))) { + break; + } + if (results.size() == 1) { + result = dyn_cast_if_present(results.front()); + if (!result) { + const auto value = cast(results.front()); + if (value != clone->getResult(0)) { + return classical.bindFrom(value, op.getResult(0), &op); + } + } + } else if (!results.empty()) { + break; } - const auto amount = static_cast(rhs->getZExtValue()); - result = isa(&op) ? lhs->shl(amount) : lhs->lshr(amount); } - bindInteger(op.getResult(0), result, classical); + if (!result) { + 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) { + 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); + arith::SubIOp, arith::MulIOp, arith::ShLIOp, arith::ShRUIOp, + arith::ShRSIOp, arith::CmpIOp, arith::AddFOp, arith::SubFOp, + arith::MulFOp, arith::DivFOp, arith::RemFOp, arith::NegFOp, + arith::CmpFOp, arith::SIToFPOp, arith::UIToFPOp>( + [&](Operation* foldable) { + return foldClassicalOp(*foldable, classical); }) - .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); - return success(); + .Case([&](arith::DivUIOp value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.udiv(rhs); + }); + }) + .Case([&](arith::DivSIOp value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.sdiv(rhs); + }); + }) + .Case([&](arith::RemUIOp value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.urem(rhs); + }); + }) + .Case([&](arith::RemSIOp value) { + return applyDivision( + value, classical, + [](const llvm::APInt& lhs, const llvm::APInt& rhs) { + return lhs.srem(rhs); + }); }) .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)) { - return failure(); - } - bindInteger(select.getResult(), *cond ? *t : *f, classical); - return success(); + 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 applyUnsignedIndexCast(cast.getIn(), cast.getOut(), cast, - classical); + return applyIntegerCast(cast.getIn(), cast.getOut(), cast, classical, + false); }) - .Case([&](arith::ExtUIOp cast) { - return applyUnsignedIndexCast(cast.getIn(), cast.getOut(), cast, - classical); + .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( + [&](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: " @@ -572,8 +950,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); @@ -600,45 +978,87 @@ 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}; } +static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, + WalkState& walk, Operation* op); + static LogicalResult bindValuePairs(ValueRange sources, ValueRange dests, WalkState& walk, Operation* op) { - const QubitMap sourceQubits = *walk.qubits; - 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 = walk.tensors->lookup(src); + if (!slots) { + return op->emitError() + << "qtensor SSA value is not mapped for QCO DD simulation"; + } + 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 = walk.classical->memrefs.find(src); + if (it == walk.classical->memrefs.end()) { + return op->emitError() + << "classical memref is not mapped for QCO DD simulation"; + } + values.emplace_back(it->second); } else { - const auto value = sourceClassical.scalars.find(src); - if (value == sourceClassical.scalars.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->scalars[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(); } +static LogicalResult bindYieldResults(YieldOp yield, + ValueRange classicalResults, + ValueRange linearResults, + WalkState& walk) { + const size_t numClassical = classicalResults.size(); + 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 +1073,202 @@ static LogicalResult walkBlock(Block& block, WalkState& walk, StateDD& state) { } template -static LogicalResult applyRegionBranch(ValueRange linearOperands, Block& block, - WalkState& walk, StateDD& state, - Operation* parent) { +static LogicalResult +applyRegionBranch(ValueRange linearOperands, Block& block, + ValueRange classicalResults, ValueRange linearResults, + WalkState& walk, StateDD& state, Operation* parent) { if (failed( bindValuePairs(linearOperands, block.getArguments(), walk, parent))) { return failure(); } + if (failed(consumeExecutionStep(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(consumeExecutionStep(walk, parent))) { + return failure(); + } + 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 (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::make_shared(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::make_shared(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 || failed(index)) { + if (!input) { + 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"; + } + 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(), input); + 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 || !wire || failed(index)) { + if (!input || !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"; + } + (*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())) { + 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 +1279,9 @@ static LogicalResult applyOp(Operation& op, WalkState& walk, StateDD& state) { .template Case([&](cbit::StoreOp store) { return storeRegister(store, *walk.classical); }) - .template Case( - [&](Operation* classicalOp) { - return applyClassicalOp(*classicalOp, *walk.classical); - }) + .template Case([](auto) { return success(); }) .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 +1307,8 @@ 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()] = + BoolAttr::get(measureOp.getContext(), bit == '1'); walk.qubits->bind(measureOp.getQubitOut(), *q); return success(); } @@ -752,84 +1339,136 @@ 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(); + 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; - } + } + 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); + 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 (failed(consumeExecutionStep(walk, forOp))) { + return failure(); + } + auto iterArgs = body.getArguments().drop_front(); + if (failed(bindValuePairs(carried, iterArgs, walk, forOp))) { + return failure(); + } + 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 { + 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 = whileOp.getConditionOp(); + auto value = + lookupBool(condition.getCondition(), *walk.classical, whileOp); + if (failed(value)) { + return failure(); + } + if (!*value) { + return bindValuePairs(condition.getArgs(), whileOp.getResults(), + walk, whileOp); + } + if (failed(consumeExecutionStep(walk, whileOp))) { + return failure(); + } + if (failed(bindValuePairs(condition.getArgs(), after.getArguments(), + walk, whileOp)) || + failed(walkBlock(after, walk, state))) { + return failure(); } - return bindValuePairs(carried, forOp.getResults(), walk, forOp); + auto yield = whileOp.getYieldOp(); + carried.assign(yield.getOperands().begin(), + yield.getOperands().end()); } }) .template Case([&](func::CallOp call) -> LogicalResult { @@ -850,6 +1489,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(); @@ -864,7 +1507,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 +1529,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 +1552,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,47 +1563,112 @@ 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; + ClassicalEnv classical; +}; +} // namespace + +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"; } - QubitMap qubits; + PreparedState prepared; + if (failed( + applyArgumentBindings(func, argumentBindings, prepared.classical))) { + return failure(); + } + 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); } if (qubits.numQubits == 0) { - qc::Qubit next = 0; + size_t next = 0; for (Value arg : func.getArguments()) { - if (!isa(arg.getType())) { - continue; + if (isa(arg.getType())) { + 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); + if (type.isDynamicDim(0)) { + const auto binding = argumentBindings.find(arg); + if (binding == argumentBindings.end()) { + return func.emitError() + << "dynamic qtensor arguments require an index extent"; + } + size = cast(binding->second).getValue().getSExtValue(); + if (size < 0) { + return func.emitError() + << "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(count); + for (size_t i = 0; i < count; ++i) { + slots.emplace_back(static_cast(next++)); + } + prepared.tensors.bind(arg, + std::make_shared(std::move(slots))); } - qubits.bind(arg, next++); } qubits.numQubits = 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()) { + 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++)); + } } 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 DDArgumentBindings& argumentBindings) { + auto prepared = + prepare(func, dd, argumentBindings, /*bindEntryAllocations=*/true); + if (failed(prepared)) { return failure(); } - QubitMap qubits = std::move(*qubitsOr); - ClassicalEnv classical; - WalkState walkState{ - .qubits = &qubits, .classical = &classical, .dd = &dd, .rng = nullptr}; + QubitMap qubits = std::move(prepared->qubits); + TensorMap tensors = std::move(prepared->tensors); + ClassicalEnv classical = std::move(prepared->classical); + WalkState walkState{.qubits = &qubits, + .tensors = &tensors, + .classical = &classical, + .dd = &dd, + .rng = nullptr}; dd::MatrixDD state = qubits.numQubits == 0 @@ -976,20 +1685,23 @@ 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 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; - ClassicalEnv classical; + QubitMap qubits = prepared.qubits; + qubits.numQubits = inputQubits; + TensorMap tensors = prepared.tensors.clone(); + ClassicalEnv classical = prepared.classical; WalkState walkState{.qubits = &qubits, + .tensors = &tensors, .classical = &classical, .dd = &dd, .rng = rng, @@ -1007,13 +1719,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 DDArgumentBindings& argumentBindings) { + auto prepared = prepare(func, dd, argumentBindings); + if (failed(prepared)) { dd.decRef(in); return failure(); } - return simulateImpl(func, in, dd, *qubits, &rng); + return simulateImpl(func, in, dd, *prepared, &rng); } static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { @@ -1026,20 +1739,28 @@ static bool isOutputOnlyRegister(Value reg, ArrayRef outputs) { }); } +static bool hasOutputOnlyMeasurementResult(MeasureOp measure, + ArrayRef outputs) { + return llvm::all_of(measure.getResult().getUses(), [&](const OpOperand& use) { + auto store = dyn_cast(use.getOwner()); + return store && isOutputOnlyRegister(store.getReg(), outputs); + }); +} + static bool isDeferrableMeasurement(MeasureOp measure, Block* entry, 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 (measure->getBlock() != entry || + !hasOutputOnlyMeasurementResult(measure, outputs)) { + return false; + } + 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, Block* entry, +static void analyzeSampling(func::FuncOp func, Block* sampledEntry, ArrayRef outputs, DenseSet& active, SamplingPlan& plan) { Operation* funcOp = func.getOperation(); @@ -1051,7 +1772,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, sampledEntry, outputs)) { plan.deferredMeasurements.insert(op); } else { plan.dynamic = true; @@ -1059,10 +1780,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, sampledEntry, outputs, active, plan); } } }); @@ -1094,7 +1816,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 +1832,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"; @@ -1123,9 +1844,10 @@ static FailureOr encodeOutcome(ArrayRef outputs, } FailureOr> -sample(func::FuncOp func, dd::Package& dd, size_t shots, std::mt19937_64& rng) { - auto qubits = prepare(func, dd); - if (failed(qubits)) { +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(); } auto plan = getSamplingPlan(func); @@ -1138,10 +1860,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(); } @@ -1152,7 +1874,7 @@ 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, + simulateImpl(func, dd::makeZeroState(numQubits, dd), dd, *prepared, nullptr, &plan->deferredMeasurements, &classical); if (failed(state)) { return failure(); @@ -1169,7 +1891,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, 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..5210e33d5e 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_dd_functionality.cpp @@ -21,11 +21,14 @@ #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 #include @@ -62,8 +65,9 @@ class QCODDFunctionalityTest : public testing::Test { void SetUp() override { DialectRegistry registry; - registry.insert(); + registry.insert(); context = std::make_unique(); context->appendDialectRegistry(registry); context->loadAllAvailableDialects(); @@ -513,6 +517,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)); @@ -564,6 +580,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 +682,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()); - - const auto elseOut = - simulate(mainFunc(*elseMod), dd::makeZeroState(1, *dd), *dd, rng); - ASSERT_TRUE(succeeded(elseOut)); - EXPECT_EQ(elseOut->getVector(), zero.getVector()); + qc::QuantumComputation thenQc(1); + thenQc.x(0); + expectEqualToQc(mainFunc(*thenMod), thenQc); - 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 +714,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) { @@ -941,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; }}, @@ -1137,6 +1156,59 @@ 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"}) { + expectMlirSimulationFails(1, source); + } +} TEST_F(QCODDFunctionalityTest, RejectsUnmappedClassicalControl) { for (const StringRef source : {R"mlir( module { @@ -1254,6 +1326,40 @@ 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"}) { + expectMlirSimulationFails(1, source); + } +} TEST_F(QCODDFunctionalityTest, Rejects) { { auto mod = buildModule([](QCOProgramBuilder& b) { @@ -1596,6 +1702,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); + DDArgumentBindings 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 { @@ -1748,6 +1875,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); @@ -1853,6 +2013,23 @@ 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); + 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 +2092,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 +2211,622 @@ 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); + DDArgumentBindings 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); + + 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, 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); + 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, 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()); + 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); + 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); + DDArgumentBindings bindings; + bindings[func.getArgument(0)] = + IntegerAttr::get(IndexType::get(context.get()), 2); + + auto dd = std::make_unique(2); + 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))); + bindings[func.getArgument(0)] = + IntegerAttr::get(IntegerType::get(context.get(), 64), 2); + EXPECT_TRUE(failed(buildFunctionality(func, *dd, bindings))); +} + +TEST_F(QCODDFunctionalityTest, RejectsQTensorBeyondQubitRange) { + expectMlirFails(1, R"mlir( + module { + func.func @main(%qubits: tensor<65537x!qco.qubit>) { + return + } + } + )mlir"); +} + +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); + 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 + %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 + %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 + %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 + %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 + %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 + %all_integers_ok = arith.andi %integer_and_casts, %signed_remainder_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) { + %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); +} + +TEST_F(QCODDFunctionalityTest, BindingsDriveObservableClassicalPath) { + auto mod = parseSourceString(R"mlir( + module { + 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 + %integer_ok = arith.cmpi eq, %sum, %expected : i16 + %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, %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 + 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); + + auto func = mainFunc(*mod); + DDArgumentBindings bindings; + bindings[func.getArgument(0)] = + IntegerAttr::get(IndexType::get(context.get()), 3); + bindings[func.getArgument(1)] = + 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(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, 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 { + 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 = arith.constant 0 : index + %reg = memref.alloc() : memref<1xi16> + %value = memref.load %reg[%index] : 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 + %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", + 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(%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(); + DDArgumentBindings 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()); + expectMlirFails(1, 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"); + expectMlirFails(1, R"mlir( + module { + func.func @main() { + %one = arith.constant 1 : index + %tensor = qtensor.alloc(%one) : tensor + qtensor.dealloc %tensor : tensor + return + } + } + )mlir"); + ASSERT_TRUE(topLevel); + + auto dd = std::make_unique(1); + const auto functionality = buildFunctionality(mainFunc(*topLevel), *dd); + ASSERT_TRUE(succeeded(functionality)); + dd->decRef(*functionality); +} + } // namespace diff --git a/test/python/test_qco_dd.py b/test/python/test_qco_dd.py index 2aa732115e..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."""