diff --git a/.agent/plans/qc-function-model.md b/.agent/plans/qc-function-model.md new file mode 100644 index 0000000000..a792178296 --- /dev/null +++ b/.agent/plans/qc-function-model.md @@ -0,0 +1,215 @@ +# Add reusable QC functions and unitary calls + +This ExecPlan is a living document. The sections `Progress`, +`Surprises & Discoveries`, `Decision Log`, and `Outcomes & Retrospective` must +be kept up to date as work proceeds. + +This ExecPlan must be maintained in accordance with `.agent/PLANS.md` from the +repository root. + +## Purpose / Big Picture + +MQT Core currently represents an imported quantum program as one QC function. +After this change, a frontend can preserve a reusable helper as a private +`func.func`, call a generic helper with `func.call`, and call a gate definition +with `qc.call`. A `qc.call` is a unitary operation, so the existing QC modifier +and analysis code can handle a custom gate without expanding its body. + +The focused QC IR unit tests demonstrate the feature by building a generic +helper and a marked unitary helper with `QCProgramBuilder`, verifying the +module, and checking that the builder selected the correct call operation. + +## Progress + +- [x] (2026-09-02 21:17Z) Inspected the MQT metadata dialect, QC operation + interfaces, QC builder state, and current QC unit tests. +- [x] (2026-09-02 22:05Z) Added the frontend-neutral unitary function marker and + its QC verifier. +- [x] (2026-09-02 22:05Z) Added `qc.call` as a unitary, symbol-using call + operation. +- [x] (2026-09-02 22:05Z) Added callback-complete QC builder APIs for generic + and unitary functions. +- [x] (2026-09-02 22:08Z) Added and passed focused marker, call, modifier, and + builder tests; all 344 QC IR unit tests pass. +- [x] (2026-09-02 23:33Z) Added PR #2336 to the existing general compiler launch + changelog entry. +- [x] (2026-09-02) Ran all focused QC tests, a fresh release build, and lint + before the final rebase. +- [x] (2026-09-02 22:55Z) Applied the independent MLIR/C++ specialist review: + static qubits are function-local, allocation mode remains module-wide, and + builder calls validate module ownership and exact operand types. +- [x] (2026-09-02 23:25Z) Applied the specialist's final verifier correction: a + unitary QC function must end in an empty `func.return`; the final delta + review found no remaining actionable issues. +- [x] (2026-09-03 00:17Z) Rebased onto `origin/main` after PR #2337 fixed the + multiplexer benchmark. The release build, all 3,805 runnable repository + tests (3,806 registered, one expected skip), and the full lint pass. + +## Surprises & Discoveries + +- Observation: The MQT dialect already links the QC and QCO dialect libraries + and already verifies operation, function-argument, and function-result + metadata. Evidence: `mlir/lib/Dialect/MQT/IR/CMakeLists.txt` and + `MQTDialect::verifyOperationAttribute` provide the required ownership point + without a new dialect or library. +- Observation: The existing `build/release` cache still referenced MLIR 22 and + failed in unrelated current-main APIs. Evidence: configuring the same preset + with `--fresh` selected the repository-configured MLIR 23.1.0 installation, + after which the QC target built successfully. +- Observation: Restoring `allocationMode` after a helper callback allowed one + module to mix static and dynamic allocation even though both conversions make + that choice once per module. Function-local allocation caches must therefore + remain separate while the mode itself remains shared. +- Observation: The multiplexer benchmark on the previous base passed a removed + `DDPackage` argument to `QCOProgram.sample`. PR #2337 fixed the benchmark on + `main` before this branch's final rebase, restoring a clean lint baseline. + +## Decision Log + +- Decision: Keep all definitions as private `func.func` operations and mark + unitary definitions with the discardable `mqt.unitary` unit attribute. + Rationale: `func.func` already supplies MLIR symbol and callable behavior; + another function operation would duplicate it. Date/Author: 2026-09-02, Codex. +- Decision: Implement unitary behavior on `qc.call`, not on `func.func`. + Rationale: MLIR operation interfaces are fixed by operation class, while one + `func.func` class must represent both generic and unitary functions. + Date/Author: 2026-09-02, Codex. +- Decision: A QC unitary function accepts zero or more `f64` parameters followed + by one or more scalar `!qc.qubit` arguments and returns no values. Its body + contains only pure scalar computations, QC unitary operations, and an empty + `func.return`. Rationale: This is the common subset required by OpenQASM gate + definitions and Qiskit `Gate` objects. Date/Author: 2026-09-02, Codex. +- Decision: Use complete builder callbacks and a `func::FuncOp` handle at call + sites. Rationale: This restores insertion state automatically, infers result + types from the completed body, and avoids paired start/end calls and + string-only symbol references. Date/Author: 2026-09-02, Codex. +- Decision: Validate a builder call against a function in the same module before + constructing IR. Rationale: the handle-based API can reject foreign symbols + and signature mismatches immediately, matching the QCO builder contract. + Date/Author: 2026-09-02, Codex and independent specialist review. +- Decision: Borrowed scalar qubit arguments are updated in place in QC and must + not also be returned explicitly. Rationale: QC-to-QCO appends their current + values to the positional ABI; allowing an explicit return would duplicate a + linear quantum value. Date/Author: 2026-09-02, Codex and independent + specialist review. + +## Outcomes & Retrospective + +QC now represents reusable generic functions with `func.call` and unitary gate +definitions with the small `mqt.unitary` plus `qc.call` contract. The builder +uses complete callbacks, so helper construction cannot leak insertion or +allocation state into the entry point. No new function operation, symbol +abstraction, or dependency was needed. The implementation passes the complete +test suite in a fresh release build. An independent specialist judged this a +strong, idiomatic MLIR 23/C++ foundation after the local-state and validation +corrections above. The specialist's final delta review found no remaining +actionable findings and judged the two-commit implementation an idiomatic MLIR +23/C++20 base for the format integrations. + +## Context and Orientation + +`mlir/include/mlir/Dialect/MQT/IR/MQTDialect.td` declares frontend-neutral +discardable attributes. `mlir/lib/Dialect/MQT/IR/MQTDialect.cpp` verifies those +attributes. `mlir/include/mlir/Dialect/QC/IR/QCInterfaces.td` defines +`qc::UnitaryOpInterface`; modifier operations and frontend exporters use this +interface to recognize unitary operations. +`mlir/include/mlir/Dialect/QC/IR/QCOps.td` defines QC operations. +`mlir/include/mlir/Dialect/QC/Builder/QCProgramBuilder.h` and its implementation +build complete QC modules and track allocations in the current function. + +The new `mqt.unitary` marker classifies a function definition. The new `qc.call` +operation refers to such a definition and implements both MLIR's +`CallOpInterface` and QC's `UnitaryOpInterface`. A generic function remains a +normal `func.func` called with `func.call`. + +## Plan of Work + +Extend the MQT dialect's discardable attributes with `mqt.unitary`. Add inline +query support and an out-of-line setter beside the existing entry-point helpers. +The MQT verifier must accept the attribute only as a unit attribute on a +private, defined, non-entry `func.func`. For a QC signature, require `f64` +parameters before scalar qubits and no results. Walk the body and accept only +regionless, memory-effect-free scalar operations, QC unitary operations, and an +empty function return. + +Define `qc.call` in `QCOps.td` with a symbol reference and a variadic operand +list. Implement `CallOpInterface`, `SymbolUserOpInterface`, and +`qc::UnitaryOpInterface`. The symbol-use verifier resolves a private marked +function and checks the exact operand signature. The unitary interface treats +all trailing qubit operands as targets and all leading operands as parameters. + +Add `createFunction`, `createUnitaryFunction`, and `call` to `QCProgramBuilder`. +Each creation method inserts one complete private helper before the entry +function under an insertion guard. It swaps the per-function allocation caches +for the callback while retaining the module-wide allocation mode, emits +deallocations for local values that are not returned, sets the inferred function +result types, emits `func.return`, and restores the entry-function state. `call` +emits `qc.call` for a marked function and `func.call` otherwise. + +Add focused tests to `mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp`. The tests +must verify valid and invalid unitary markers, symbol/signature checking, +modifier nesting, insertion restoration, generic result inference, and call +selection. Append the eventual pull request reference and contributors to the +existing general compiler launch changelog entry. Do not create another +changelog bullet. + +## Concrete Steps + +Run all commands from the repository root. + +Build and run the focused test while iterating: + + cmake --build --preset release --target mqt-core-mlir-unittest-qc-ir + ./build/release/mlir/unittests/Dialect/QC/IR/mqt-core-mlir-unittest-qc-ir + +Run final validation: + + cmake --preset release + cmake --build --preset release + ctest --preset release + uvx nox -s lint + +The focused binary and CTest must report no failures. Lint must finish without +modifying tracked files; if it formats the implementation, inspect the edits and +rerun the affected checks. + +## Validation and Acceptance + +A parsed private marked QC helper with only unitary operations verifies. A +marked entry point, declaration, result-bearing function, non-`f64` parameter, +non-qubit trailing argument, allocation, measurement, reset, or recursive call +does not verify. + +`QCProgramBuilder::createFunction` returns a private function whose result types +match its callback results. `createUnitaryFunction` returns a marked resultless +function. `call` emits `func.call` for the first and `qc.call` for the second. +After either creation callback, subsequent entry operations remain in `main`. A +`qc.call` can appear in a QC inverse or control modifier because it implements +`qc::UnitaryOpInterface`. + +## Idempotence and Recovery + +All builds and tests are repeatable. The implementation is the first commit of +one self-contained branch from `origin/main`, followed by the QCO work in its +companion ExecPlan. If an edit fails, inspect `git diff` and use another focused +patch; do not reset or discard unrelated work. + +## Artifacts and Notes + +The repository was clean before the branch was created. The branch starts at the +current `origin/main` commit. + +## Interfaces and Dependencies + +The public C++ builder interface will contain: + + func::FuncOp createFunction( + StringRef name, TypeRange argumentTypes, + function_ref(ValueRange)> body); + func::FuncOp createUnitaryFunction( + StringRef name, TypeRange argumentTypes, + function_ref body); + SmallVector call(func::FuncOp callee, ValueRange operands); + +No new dependency is required. The implementation uses the existing Func, MQT, +QC, and MLIR call/symbol interfaces. diff --git a/.agent/plans/qco-function-model.md b/.agent/plans/qco-function-model.md new file mode 100644 index 0000000000..229afc522c --- /dev/null +++ b/.agent/plans/qco-function-model.md @@ -0,0 +1,179 @@ +# Add value-semantic functions and calls to QCO + +This ExecPlan is a living document maintained according to `.agent/PLANS.md`. + +## Purpose / Big Picture + +QCO needs a direct representation for reusable unitary definitions and a +loss-minimizing convention for ordinary functions that thread qubits through SSA +results. After this change, `qco.call` is an ordinary unitary operation, generic +`func.call` is an explicit wire boundary, and QC/QCO conversion uses one +positional function ABI instead of per-result annotations or deriving +interprocedural correspondence by walking callee bodies. + +## Progress + +- [x] (2026-09-02 21:36Z) Compared current main, PR #2196, the QCO builder, + WireIterator, and both conversion passes. +- [x] (2026-09-02 22:28Z) Chose a positional QCO function ABI and removed the + proposed result annotation. +- [x] (2026-09-02 22:28Z) Added `qco.call` and QCO unitary-function + verification. +- [x] (2026-09-02 22:28Z) Added callback-complete QCO builder function and call + APIs. +- [x] (2026-09-02 22:28Z) Made both QC/QCO conversions preserve scalar-qubit + functions and calls. +- [x] (2026-09-02 22:28Z) Removed generic-call inference and caching from the + wire and tensor iterators. +- [x] (2026-09-02) Added focused generic-call and QC/QCO round-trip tests. +- [x] (2026-09-02 23:33Z) Added PR #2336 to the existing general compiler launch + changelog entry. +- [x] (2026-09-02) Ran 1,157 focused QC, QCO, and conversion tests and lint + before the final rebase. +- [x] (2026-09-02 22:55Z) Applied the independent MLIR/C++ specialist review: + malformed calls fail safely, call metadata round-trips or fails explicitly + when QC cannot represent it, and builder/call contracts match across the + two dialects. +- [x] (2026-09-02 23:25Z) Applied the specialist's final corrections: QC-to-QCO + rejects duplicate or explicitly returned borrowed qubits before mutation, + QCO-to-QC rejects attributes it cannot preserve on stripped pass-through + results, and no redundant nested call verifier remains. The final delta + review found no remaining actionable issues. +- [x] (2026-09-03 00:17Z) Rebased onto `origin/main` after PR #2337 fixed the + multiplexer benchmark. The release build, all 3,805 runnable repository + tests (3,806 registered, one expected skip), and the full lint pass. + +## Surprises & Discoveries + +- PR #2196 adds 212 builder implementation lines largely because it exposes + paired start/end state and recomputes qubit and tensor correspondence from + callee bodies. A positional ABI and `qco.call` make both responsibilities + local and remove those failure modes. +- Current QC-to-QCO preflight rejects function qubit block arguments before + dialect conversion starts. Multi-function support therefore requires an + explicit function conversion path; changing only `func.call` is insufficient. +- MLIR function signature conversion temporarily creates type-conversion casts + that do not satisfy the final unitary-function signature. Both conversion + passes therefore hide the marker under a scope guard and restore it after the + whole conversion, including failure paths. +- Converted QC block arguments must retain the original QC SSA value as their + state-map key. Treating the converted QCO argument as a second key returns the + stale argument instead of the latest qubit at `func.return`. +- The existing `build/release` directory still contains generated Neutral Atom + QDMI manifests from an older checkout. They add a fifth discovered device and + make two unrelated registry tests fail; final validation therefore uses a + fresh build directory rather than deleting the user's existing build state. +- A malformed `qco.call` was parsed far enough for the function attribute + verifier to query its unitary interface before the operation verifier ran. Its + correspondence accessors and enclosing verifier must therefore be total even + on invalid IR rather than relying on assertions or verifier order. +- `func.call`, `qc.call`, and `qco.call` carry argument, result, and discardable + attributes. Conversion must preserve all representable metadata and reject + nonempty attributes on synthetic QCO qubit results instead of silently + dropping them when converting to resultless QC calls. + +## Decision Log + +- Decision: A QCO function places source-language results first, followed by one + updated qubit for every scalar qubit argument in qubit-argument order. + QCO-to-QC validates the correspondence before stripping those trailing values. + No result annotation is used. Rationale: the formats being targeted borrow + fixed qubit operands rather than returning arbitrary qubit identities; one + positional convention is explicit, loss-minimizing, and cannot become stale + independently of the signature. Date/Author: 2026-09-02, user and Codex. +- Decision: A marked QCO unitary function has `f64` parameters followed by + qubits and returns those qubits positionally. `qco.call` has the same direct + input/output mapping and an unknown compile-time matrix. Rationale: this is + enough for gate definitions, modifiers, WireIterator, and format frontends + without inlining or matrix synthesis. Date/Author: 2026-09-02, Codex. +- Decision: Generic `func.call` is a WireIterator boundary. Rationale: generic + functions may measure, reset, allocate, branch, or return unrelated qubits; + interprocedural consumers use the positional ABI directly rather than making + local wire iteration infer whole-callee behavior. Date/Author: 2026-09-02, + Codex. +- Decision: Builder APIs take complete callbacks and function handles. + Rationale: insertion and linear-tracking state cannot leak across a paired + start/end API, while result types are inferred once from the completed body. + Date/Author: 2026-09-02, Codex. +- Decision: Keep malformed-call safety and metadata preservation local to the + call operations and conversion patterns. Rationale: these are trust-boundary + correctness checks; another ABI descriptor or annotation layer would duplicate + the positional convention. Date/Author: 2026-09-02, Codex and independent + specialist review. + +## Context and Orientation + +`mlir/include/mlir/Dialect/QCO/IR/QCOOps.td` defines value-semantic quantum +operations and `qco::UnitaryOpInterface`. `QCOProgramBuilder` tracks live qubit +SSA values. Before this change, `WireIterator` followed generic calls through a +cached `CallQubitMapping`. `QCToQCO.cpp` and `QCOToQC.cpp` currently rely on +MLIR's type-only function/call conversion and therefore do not add or strip +qubit results. + +## Plan of Work + +Extend the unitary marker verifier to accept the QCO signature and prove each +returned qubit traces back to the corresponding argument through QCO unitary +operations. + +Define `qco.call` with call/symbol interfaces and the QCO unitary interface. Its +qubit inputs and outputs correspond positionally. Add complete callback builder +APIs, validate the trailing positional results from local QCO wire flow, and +update live-value tracking from the function signature at generic calls. + +Teach QC-to-QCO to append the latest value of each qubit function argument to +the function return and to convert `qc.call` to `qco.call`. Teach QCO-to-QC to +validate and strip those trailing pass-through qubit results from function +signatures, returns, and call sites, replacing each stripped call result with +the corresponding operand. Earlier results remain ordinary converted results. +Preserve call attributes in both directions; reject result attributes attached +to QCO-only pass-through qubits because QC has nowhere to store them. + +Delete `CallQubitMapping`, its cache/invalidation API, and the special +`func.call` branches from WireIterator. `qco.call` needs no special iterator +code because it implements `UnitaryOpInterface`. + +## Concrete Steps + +Run focused builds and tests while iterating: + + cmake --build --preset release --target mqt-core-mlir-unittest-qco-ir mqt-core-mlir-unittest-qco-utils mqt-core-mlir-unittest-qc-to-qco mqt-core-mlir-unittest-qco-to-qc + ctest --test-dir build/release -R 'QCO|QCToQCO|QCOToQC' --output-on-failure + +Run final validation: + + cmake --build --preset release + ctest --preset release + uvx nox -s lint + +## Validation and Acceptance + +A marked QCO helper verifies and can be nested under QCO modifiers. Its call is +traversed by WireIterator through the unitary interface. A generic call ends +wire traversal. QC-to-QCO followed by QCO-to-QC preserves a helper call and does +not leave redundant pass-through qubit results in QC. QCO-to-QC followed by +QC-to-QCO reconstructs the same positional ABI for the supported one-block outer +function shape. + +## Idempotence and Recovery + +The work is isolated on `codex/qco-function-model` as the second commit of one +self-contained branch from `origin/main`. It will be published as one new PR, +independent of PR #2196 and its stack. Builds and tests are repeatable. + +## Outcomes & Retrospective + +The positional ABI supports generic and unitary scalar-qubit functions in the +QCO builder and in both conversion directions. `qco.call` gives local quantum +analyses an explicit unitary edge, while generic calls deliberately stop local +wire traversal. Removing the speculative generic-call inference deleted both +mapping caches and their failure-prone body analysis. The implementation passes +the complete test suite in a fresh release build; the existing release build +remains contaminated by obsolete generated QDMI manifests and was left intact. +After the specialist corrections, the call accessors are safe on malformed IR, +and call metadata is either preserved losslessly or rejected when QC cannot +represent it. No result annotations, recursive matrix synthesis, or generic-call +mapping abstraction was added. The specialist's final delta review found no +remaining actionable findings and judged the positional ABI and local iterator +boundary an idiomatic MLIR 23/C++20 foundation for OpenQASM, Qiskit, and jeff +integration. diff --git a/AGENTS.md b/AGENTS.md index 9e173715ac..a57fb7147e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -116,7 +116,11 @@ MQT Core. The project-wide policy for AI-assisted contributions is The C++ code targets C++20 and uses GoogleTest. Follow these rules: -- Write Doxygen comments with `///`. +- Write Doxygen API and `@file` descriptions with `///`, preserving their + content. Keep `//!<` or `///<` for trailing member documentation and block + documentation inside continued macros. +- Use `//` for ordinary code comments and namespace closing comments. Keep + inline `/* ... */` comments, including unused parameter names. - Use `#pragma once` in headers and use existing project abstractions. - Prefer C++20 standard-library facilities over custom equivalents. - Within the `mlir` namespace and its nested namespaces, prefer LLVM types such @@ -128,6 +132,7 @@ The C++ code targets C++20 and uses GoogleTest. Follow these rules: the header that provides each type. - Do not use `module` as a C++ variable or parameter name because it conflicts with the C++20 keyword. Use `moduleOp` for `mlir::ModuleOp` values. +- Generally give non-public data members a trailing underscore. - Follow the canonical general and MLIR-specific coding policies in [`docs/development.md`](docs/development.md) and [`docs/mlir/development.md`](docs/mlir/development.md). diff --git a/CHANGELOG.md b/CHANGELOG.md index 283327b8f7..a37677325b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,9 +26,9 @@ releases may include breaking changes. [#1927], [#1935], [#1936], [#1938], [#1975], [#1976], [#2006], [#2014], [#2015], [#2017], [#2026], [#2028], [#2054], [#2058], [#2125], [#2136], [#2149], [#2150], [#2158], [#2194], [#2210], [#2211], [#2215], [#2218], - [#2220], [#2323]) ([**@burgholzer**], [**@denialhaag**], [**@taminob**], - [**@DRovara**], [**@li-mingbao**], [**@Ectras**], [**@MatthiasReumann**], - [**@simon1hofmann**], [**@J4MMlE**]) + [#2220], [#2323], [#2336]) ([**@burgholzer**], [**@denialhaag**], + [**@taminob**], [**@DRovara**], [**@li-mingbao**], [**@Ectras**], + [**@MatthiasReumann**], [**@simon1hofmann**], [**@J4MMlE**]) - ✨ Add a library for typed structured quantum benchmarks with versioned instance specifications, analytic references, deterministic manifests, and C++, Python, and command-line interfaces ([#2135], [#2299], [#2315], [#2337]) @@ -870,6 +870,7 @@ for previous changelogs._ [#2337]: https://github.com/munich-quantum-toolkit/core/pull/2337 +[#2336]: https://github.com/munich-quantum-toolkit/core/pull/2336 [#2335]: https://github.com/munich-quantum-toolkit/core/pull/2335 [#2334]: https://github.com/munich-quantum-toolkit/core/pull/2334 [#2323]: https://github.com/munich-quantum-toolkit/core/pull/2323 diff --git a/docs/development.md b/docs/development.md index d441ffc5e2..608b23db29 100644 --- a/docs/development.md +++ b/docs/development.md @@ -46,11 +46,25 @@ differences that apply to code built on LLVM and MLIR. ### C++ documentation comments -Use `///` for Doxygen documentation comments. Do not use `/** ... */`. The first -sentence is the summary; separate additional paragraphs with a blank `///` line -instead of using `\brief` or `\details`. Document parameters and return values -only when the explanation adds information that the name and signature do not -already provide. +Use `///` for Doxygen documentation comments. The first sentence is the summary; +separate additional paragraphs with a blank `///` line instead of using `\brief` +or `\details`. Document parameters and return values only when the explanation +adds information that the name and signature do not already provide. Preserve +existing documentation when changing comment style. + +Keep `//!<` or `///<` for trailing member documentation. Keep `/** ... */` +documentation inside backslash-continued macros: line comments there can consume +the following declarations after line splicing. Preserve explicit `@brief` +commands there when Doxygen needs them to retain summaries after macro +expansion. + +Keep top-level `@file` documentation and put its summary on the next line, +without `@brief`: + +```cpp +/// @file Circuit.h +/// Defines the circuit representation. +``` ```cpp /// Returns the number of qubits in the circuit. @@ -71,6 +85,10 @@ Keep public API documentation in the declaration and do not duplicate it in the implementation. Use ordinary implementation comments for details that do not belong to the API contract. +Use `//` for ordinary implementation and namespace closing comments. Inline +`/* ... */` comments remain valid, including unused parameter names such as +`OpAdaptor /*adaptor*/` and argument labels such as `/*isSigned=*/false`. + ### Reproduce C++ lint locally Before pushing a C++ change, run: diff --git a/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.h b/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.h index 579dfec92a..babc0d1da7 100644 --- a/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.h +++ b/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.h @@ -35,6 +35,15 @@ void setEntryPoint(Operation* operation); /// Remove the program entry-point marker from an operation. void removeEntryPoint(Operation* operation); +/// Return whether an operation defines a unitary function. +[[nodiscard]] inline bool isUnitaryFunction(Operation* operation) { + return operation != nullptr && + operation->hasAttr(MQTDialect::UnitaryAttrHelper::getNameStr()); +} + +/// Mark a function as unitary. +void setUnitaryFunction(Operation* operation); + /// Return the program entry point, or null if the module has none. [[nodiscard]] inline func::FuncOp getEntryPoint(ModuleOp moduleOp) { for (auto function : moduleOp.getOps()) { diff --git a/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.td b/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.td index f5c762b6bc..35e686edbb 100644 --- a/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.td +++ b/mlir/include/mlir/Dialect/MQT/IR/MQTDialect.td @@ -24,6 +24,8 @@ def MQTDialect : Dialect { across quantum dialect conversions. It defines no operations or types. `mqt.input_name` records the source-level name of a function input. + `mqt.source_name` records a source-level function name when the IR symbol + must be uniquified. `mqt.parameter_group` optionally preserves the source-level vector identity, name, element index, and size of a function input or a lexically bound `scf.for` parameter. @@ -32,12 +34,15 @@ def MQTDialect : Dialect { namespace. `mqt.entry_point` marks the single defined program entry function in a module. + `mqt.unitary` marks a private function that defines a unitary operation. `#mqt.compilation_target` records compiler-target facts as typed IR. }]; let discardableAttrs = (ins "::mlir::StringAttr":$input_name, + "::mlir::StringAttr":$source_name, "::mlir::DictionaryAttr":$parameter_group, - "::mlir::StringAttr":$register_name, "::mlir::UnitAttr":$entry_point); + "::mlir::StringAttr":$register_name, "::mlir::UnitAttr":$entry_point, + "::mlir::UnitAttr":$unitary); let hasOperationAttrVerify = 1; let hasRegionArgAttrVerify = 1; diff --git a/mlir/include/mlir/Dialect/QC/Builder/QCProgramBuilder.h b/mlir/include/mlir/Dialect/QC/Builder/QCProgramBuilder.h index 9e53c58fa6..17145aa9da 100644 --- a/mlir/include/mlir/Dialect/QC/Builder/QCProgramBuilder.h +++ b/mlir/include/mlir/Dialect/QC/Builder/QCProgramBuilder.h @@ -14,6 +14,7 @@ #include #include +#include #include #include #include @@ -34,73 +35,79 @@ class ValueRange; namespace qc { -/** - * @brief Builder API for constructing quantum programs in the QC dialect - * - * @details - * The QCProgramBuilder provides a type-safe interface for constructing - * quantum circuits using reference semantics. Operations modify qubits in - * place without producing new SSA values, providing a natural mapping to - * hardware execution models. - * - * @par Qubit addressing: - * A program must use either static qubits (`staticQubit`) or dynamic allocation - * (`allocQubit` / `allocQubitRegister`), never both. The builder terminates - * with a usage error if the modes are mixed. - * - * @par Example Usage: - * ```c++ - * QCProgramBuilder builder(context); - * builder.initialize(); - * - * auto q0 = builder.staticQubit(0); - * auto q1 = builder.staticQubit(1); - * - * // Operations modify qubits in place - * builder.h(q0).cx(q0, q1); - * - * auto module = builder.finalize(); - * ``` - */ +/// Builder API for constructing quantum programs in the QC dialect +/// +/// The QCProgramBuilder provides a type-safe interface for constructing +/// quantum circuits using reference semantics. Operations modify qubits in +/// place without producing new SSA values, providing a natural mapping to +/// hardware execution models. +/// +/// @par Qubit addressing: +/// A program must use either static qubits (`staticQubit`) or dynamic +/// allocation +/// (`allocQubit` / `allocQubitRegister`), never both. The builder terminates +/// with a usage error if the modes are mixed. +/// +/// @par Example Usage: +/// ```c++ +/// QCProgramBuilder builder(context); +/// builder.initialize(); +/// +/// auto q0 = builder.staticQubit(0); +/// auto q1 = builder.staticQubit(1); +/// +/// // Operations modify qubits in place +/// builder.h(q0).cx(q0, q1); +/// +/// auto moduleOp = builder.finalize(); +/// ``` class QCProgramBuilder final : public ImplicitLocOpBuilder { public: - /** - * @brief Construct a new QCProgramBuilder - * @param context The MLIR context to use for building operations - */ + /// Construct a new QCProgramBuilder + /// @param context The MLIR context to use for building operations explicit QCProgramBuilder(MLIRContext* context); //===--------------------------------------------------------------------===// // Initialization //===--------------------------------------------------------------------===// - /** - * @brief Initialize the builder and prepare for program construction, with - * a default return type of i64. - * - * @details - * Creates a main function with an `mqt.entry_point` attribute. Must be called - * before adding operations. - */ + /// Initialize the builder and prepare for program construction, with + /// a default return type of i64. + /// + /// Creates a main function with an `mqt.entry_point` attribute. Must be + /// called before adding operations. void initialize(); - /** - * @brief Initialize the builder and prepare for program construction - * with specified return types. - * @param returnTypes The return types for the main function - * - * @details - * Creates a main function with an `mqt.entry_point` attribute. Must be called - * before adding operations. - */ + /// Initialize the builder and prepare for program construction + /// with specified return types. + /// @param returnTypes The return types for the main function + /// + /// Creates a main function with an `mqt.entry_point` attribute. Must be + /// called before adding operations. void initialize(TypeRange returnTypes); - /** - * @brief Modify the return types of the main function after initialization. - * @param returnTypes The new return types for the main function - */ + /// Modify the return types of the main function after initialization. + /// @param returnTypes The new return types for the main function void retype(TypeRange returnTypes); + //===--------------------------------------------------------------------===// + // Functions + //===--------------------------------------------------------------------===// + + /// Create a complete private function and infer its result types. + /// + /// Borrowed qubit arguments are updated in place and must not be returned. + func::FuncOp + createFunction(StringRef name, TypeRange argumentTypes, + function_ref(ValueRange)> body); + + /// Create a complete private unitary function. + func::FuncOp createUnitaryFunction(StringRef name, TypeRange argumentTypes, + function_ref body); + + /// Call a function, using `qc.call` for a unitary function. + SmallVector call(func::FuncOp callee, ValueRange operands); + //===--------------------------------------------------------------------===// // Constants //===--------------------------------------------------------------------===// @@ -121,56 +128,46 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { // Memory Management //===--------------------------------------------------------------------===// - /** - * @brief Represents a qubit register with its qubits. - */ + /// Represents a qubit register with its qubits. struct QubitRegister { /// The memref value representing the qubit register Value value; /// The allocated qubit values SmallVector qubits; - /** - * @brief Access a specific qubit in the register - * @param index The index of the qubit to access - * @return The specified qubit value - */ + /// Access a specific qubit in the register + /// @param index The index of the qubit to access + /// @return The specified qubit value Value operator[](size_t index) const; - /** - * @brief Conversion to the backing memref value - * @return The memref value representing the qubit register - */ + /// Conversion to the backing memref value + /// @return The memref value representing the qubit register explicit operator Value() const { return value; } }; - /** - * @brief Allocate a single qubit initialized to |0⟩ - * @return A qubit reference - * - * @par Example: - * ```c++ - * auto q = builder.allocQubit(); - * ``` - * ```mlir - * %q = qc.alloc : !qc.qubit - * ``` - */ + /// Allocate a single qubit initialized to |0⟩ + /// @return A qubit reference + /// + /// @par Example: + /// ```c++ + /// auto q = builder.allocQubit(); + /// ``` + /// ```mlir + /// %q = qc.alloc : !qc.qubit + /// ``` Value allocQubit(); - /** - * @brief Get a static qubit by index - * @param index The qubit index - * @return A qubit reference - * - * @par Example: - * ```c++ - * auto q0 = builder.staticQubit(0); - * ``` - * ```mlir - * %q0 = qc.static 0 : !qc.qubit - * ``` - */ + /// Get a static qubit by index + /// @param index The qubit index + /// @return A qubit reference + /// + /// @par Example: + /// ```c++ + /// auto q0 = builder.staticQubit(0); + /// ``` + /// ```mlir + /// %q0 = qc.static 0 : !qc.qubit + /// ``` Value staticQubit(uint64_t index); /// Allocate a qubit register and eagerly load every element. @@ -203,44 +200,40 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { /// \returns The memref value that represents the qubit register. Value allocQubitRegisterStorage(int64_t size, StringRef name = {}); - /** - * @brief Explicitly loads a qubit from a memref - * - * @param memref Source memref - * @param index The index from where the qubit is loaded - * @return The loaded qubit - * - * @par Example: - * ```c++ - * auto q0 = builder.loadQubit(memref, index); - * ``` - * ```mlir - * %q0 = memref.load %memref[%index] : memref<3x!qc.qubit> - * ``` - */ + /// Explicitly loads a qubit from a memref + /// + /// @param memref Source memref + /// @param index The index from where the qubit is loaded + /// @return The loaded qubit + /// + /// @par Example: + /// ```c++ + /// auto q0 = builder.loadQubit(memref, index); + /// ``` + /// ```mlir + /// %q0 = memref.load %memref[%index] : memref<3x!qc.qubit> + /// ``` Value loadQubit(Value memref, Value index); - /** - * @brief Allocate a classical bit register - * - * @details The register uses `!cbit.reg`. Its initialization is explicit - * and independent of every other register built by this builder. - * - * @param size Number of bits (must be positive) - * @param name Optional source-level register name; defaults to no name - * @param initialization Initial value of the register elements; defaults to - * zero - * @return The CBit register value - * - * @par Example: - * ```c++ - * auto c = builder.allocClassicalBitRegister(3, "c"); - * ``` - * ```mlir - * %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} - * : !cbit.reg<3> - * ``` - */ + /// Allocate a classical bit register + /// + /// The register uses `!cbit.reg`. Its initialization is explicit + /// and independent of every other register built by this builder. + /// + /// @param size Number of bits (must be positive) + /// @param name Optional source-level register name; defaults to no name + /// @param initialization Initial value of the register elements; defaults to + /// zero + /// @return The CBit register value + /// + /// @par Example: + /// ```c++ + /// auto c = builder.allocClassicalBitRegister(3, "c"); + /// ``` + /// ```mlir + /// %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} + /// : !cbit.reg<3> + /// ``` Value allocClassicalBitRegister( int64_t size, StringRef name = {}, cbit::Initialization initialization = cbit::Initialization::Zero); @@ -256,45 +249,41 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { // Measurement and Reset //===--------------------------------------------------------------------===// - /** - * @brief Measure a qubit in the computational basis - * - * @details Measures a qubit in place and returns the classical measurement - * result. - * - * @param qubit The qubit to measure - * @return Classical measurement result (`i1`) - * - * @par Example: - * ```c++ - * auto result = builder.measure(q); - * ``` - * ```mlir - * %result = qc.measure %q : !qc.qubit -> i1 - * ``` - */ + /// Measure a qubit in the computational basis + /// + /// Measures a qubit in place and returns the classical measurement + /// result. + /// + /// @param qubit The qubit to measure + /// @return Classical measurement result (`i1`) + /// + /// @par Example: + /// ```c++ + /// auto result = builder.measure(q); + /// ``` + /// ```mlir + /// %result = qc.measure %q : !qc.qubit -> i1 + /// ``` Value measure(Value qubit); - /** - * @brief Measure a qubit and store the result in a classical bit register - * - * @details Measures the qubit and stores the classical result in the given - * classical register at the given index, in addition to returning it. - * - * @param qubit The qubit to measure - * @param reg The CBit register - * @param index The index within the classical register - * @return Classical measurement result (`i1`) - * - * @par Example: - * ```c++ - * builder.measure(q0, c, 0); - * ``` - * ```mlir - * %r0 = qc.measure %q0 : !qc.qubit -> i1 - * cbit.store %r0, %c[%c0] : !cbit.reg<3> - * ``` - */ + /// Measure a qubit and store the result in a classical bit register + /// + /// Measures the qubit and stores the classical result in the given + /// classical register at the given index, in addition to returning it. + /// + /// @param qubit The qubit to measure + /// @param reg The CBit register + /// @param index The index within the classical register + /// @return Classical measurement result (`i1`) + /// + /// @par Example: + /// ```c++ + /// builder.measure(q0, c, 0); + /// ``` + /// ```mlir + /// %r0 = qc.measure %q0 : !qc.qubit -> i1 + /// cbit.store %r0, %c[%c0] : !cbit.reg<3> + /// ``` Value measure(Value qubit, Value reg, const std::variant& index); @@ -310,23 +299,20 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { QCProgramBuilder& measureQubitRegister(Value qubits, Value bits, int64_t size); - /** - * @brief Reset a qubit to |0⟩ state - * - * @details - * Resets a qubit to the |0⟩ state in place. - * - * @param qubit The qubit to reset - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.reset(q); - * ``` - * ```mlir - * qc.reset %q : !qc.qubit - * ``` - */ + /// Reset a qubit to |0⟩ state + /// + /// Resets a qubit to the |0⟩ state in place. + /// + /// @param qubit The qubit to reset + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.reset(q); + /// ``` + /// ```mlir + /// qc.reset %q : !qc.qubit + /// ``` QCProgramBuilder& reset(Value qubit); //===--------------------------------------------------------------------===// @@ -965,173 +951,155 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { // BarrierOp - /** - * @brief Apply a BarrierOp - * - * @param qubits Target qubits - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.barrier({q0, q1}); - * ``` - * ```mlir - * qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit - * ``` - */ + /// Apply a BarrierOp + /// + /// @param qubits Target qubits + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.barrier({q0, q1}); + /// ``` + /// ```mlir + /// qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit + /// ``` QCProgramBuilder& barrier(ValueRange qubits); - /** - * @brief Apply an explicitly represented dense unitary matrix - * - * @param qubits Target qubits, ordered from the most-significant basis bit - * to the least-significant basis bit - * @param matrix Square row-major `complex` matrix - * @return Reference to this builder for method chaining - */ + /// Apply an explicitly represented dense unitary matrix + /// + /// @param qubits Target qubits, ordered from the most-significant basis bit + /// to the least-significant basis bit + /// @param matrix Square row-major `complex` matrix + /// @return Reference to this builder for method chaining QCProgramBuilder& unitary(ValueRange qubits, DenseElementsAttr matrix); //===--------------------------------------------------------------------===// // Modifiers //===--------------------------------------------------------------------===// - /** - * @brief Apply a control modifier to a collection of gates - * - * @param controls Control qubits - * @param targets Target qubits the body operates on - * @param body Function that builds the body containing the target gates - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.ctrl(q0, q1, [&](ValueRange targets) { - * builder.x(targets[0]); - * }); - * ``` - * ```mlir - * qc.ctrl(%q0) targets(%a0 = %q1) { - * qc.x %a0 : !qc.qubit - * } : !qc.qubit - * ``` - */ + /// Apply a control modifier to a collection of gates + /// + /// @param controls Control qubits + /// @param targets Target qubits the body operates on + /// @param body Function that builds the body containing the target gates + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.ctrl(q0, q1, [&](ValueRange targets) { + /// builder.x(targets[0]); + /// }); + /// ``` + /// ```mlir + /// qc.ctrl(%q0) targets(%a0 = %q1) { + /// qc.x %a0 : !qc.qubit + /// } : !qc.qubit + /// ``` QCProgramBuilder& ctrl(ValueRange controls, ValueRange targets, const function_ref& body); - /** - * @brief Apply a control modifier with a single target and one-qubit body. - * - * @param controls Control qubits - * @param target Target qubit - * @param body Function that builds the body containing the target operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.ctrl({q0_in, q1_in}, q2_in, [&](Value target) { - * builder.x(target); - * }); - * ``` - */ + /// Apply a control modifier with a single target and one-qubit body. + /// + /// @param controls Control qubits + /// @param target Target qubit + /// @param body Function that builds the body containing the target operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.ctrl({q0_in, q1_in}, q2_in, [&](Value target) { + /// builder.x(target); + /// }); + /// ``` QCProgramBuilder& ctrl(ValueRange controls, Value target, const function_ref& body); - /** - * @brief Apply a control modifier with one control and one target. - * - * @param control Control qubit - * @param target Target qubit - * @param body Function that builds the body containing the target operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.ctrl(q0_in, q1_in, [&](Value target) { - * builder.x(target); - * }); - * ``` - */ + /// Apply a control modifier with one control and one target. + /// + /// @param control Control qubit + /// @param target Target qubit + /// @param body Function that builds the body containing the target operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.ctrl(q0_in, q1_in, [&](Value target) { + /// builder.x(target); + /// }); + /// ``` QCProgramBuilder& ctrl(Value control, Value target, const function_ref& body); - /** - * @brief Apply an inverse (i.e., adjoint) modifier to a collection of gates - * - * @param qubits The qubits the body operates on - * @param body Function that builds the body containing the gates to invert - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.inv(q0, [&](ValueRange qubits) { - * builder.h(qubits[0]); - * }); - * ``` - * ```mlir - * qc.inv (%a0 = %q0) { - * qc.s %a0 : !qc.qubit - * } - * ``` - */ + /// Apply an inverse (i.e., adjoint) modifier to a collection of gates + /// + /// @param qubits The qubits the body operates on + /// @param body Function that builds the body containing the gates to invert + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.inv(q0, [&](ValueRange qubits) { + /// builder.h(qubits[0]); + /// }); + /// ``` + /// ```mlir + /// qc.inv (%a0 = %q0) { + /// qc.s %a0 : !qc.qubit + /// } + /// ``` QCProgramBuilder& inv(ValueRange qubits, const function_ref& body); - /** - * @brief Apply an inverse modifier on a single qubit. - * - * @param qubit Qubit involved in the operation - * @param body Function that builds the body containing the operation to - * invert - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.inv(q0_in, [&](Value qubit) { - * builder.h(qubit); - * }); - * ``` - */ + /// Apply an inverse modifier on a single qubit. + /// + /// @param qubit Qubit involved in the operation + /// @param body Function that builds the body containing the operation to + /// invert + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.inv(q0_in, [&](Value qubit) { + /// builder.h(qubit); + /// }); + /// ``` QCProgramBuilder& inv(Value qubit, const function_ref& body); - /** - * @brief Apply a power modifier to a collection of gates - * - * @param exponent The exponent to raise the operation to - * @param qubits The qubits the body operates on - * @param body Function that builds the body containing the gates to - * exponentiate - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.pow(2.0, {q0, q1}, [&](ValueRange qubits) { - * builder.swap(qubits[0], qubits[1]); - * }); - * ``` - * ```mlir - * qc.pow(%exponent) (%a0 = %q0) { - * qc.s %a0 : !qc.qubit - * } : !qc.qubit - * ``` - */ + /// Apply a power modifier to a collection of gates + /// + /// @param exponent The exponent to raise the operation to + /// @param qubits The qubits the body operates on + /// @param body Function that builds the body containing the gates to + /// exponentiate + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.pow(2.0, {q0, q1}, [&](ValueRange qubits) { + /// builder.swap(qubits[0], qubits[1]); + /// }); + /// ``` + /// ```mlir + /// qc.pow(%exponent) (%a0 = %q0) { + /// qc.s %a0 : !qc.qubit + /// } : !qc.qubit + /// ``` QCProgramBuilder& pow(const std::variant& exponent, ValueRange qubits, const function_ref& body); - /** - * @brief Apply a power modifier on a single qubit. - * - * @param exponent The exponent to raise the operation to - * @param qubit Qubit involved in the operation - * @param body Function that builds the body containing the operation to - * exponentiate - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.pow(2.0, q0, [&](Value qubit) { builder.s(qubit); }); - * ``` - */ + /// Apply a power modifier on a single qubit. + /// + /// @param exponent The exponent to raise the operation to + /// @param qubit Qubit involved in the operation + /// @param body Function that builds the body containing the operation to + /// exponentiate + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.pow(2.0, q0, [&](Value qubit) { builder.s(qubit); }); + /// ``` QCProgramBuilder& pow(const std::variant& exponent, Value qubit, const function_ref& body); @@ -1139,189 +1107,173 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { // Deallocation //===--------------------------------------------------------------------===// - /** - * @brief Explicitly deallocate a qubit - * - * @details - * Deallocates a qubit and removes it from tracking. Optional, finalize() - * automatically deallocates all remaining allocated qubits. - * - * @param qubit The qubit to deallocate - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.dealloc(q); - * ``` - * ```mlir - * qc.dealloc %q : !qc.qubit - * ``` - */ + /// Explicitly deallocate a qubit + /// + /// Deallocates a qubit and removes it from tracking. Optional, finalize() + /// automatically deallocates all remaining allocated qubits. + /// + /// @param qubit The qubit to deallocate + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.dealloc(q); + /// ``` + /// ```mlir + /// qc.dealloc %q : !qc.qubit + /// ``` QCProgramBuilder& dealloc(Value qubit); //===--------------------------------------------------------------------===// // SCF operations //===--------------------------------------------------------------------===// - /** - * @brief Construct an scf.for operation - * - * @param lowerbound Lower bound of the loop - * @param upperbound Upper bound of the loop - * @param step Step size of the loop - * @param body Function that builds the body of the for operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfFor(lb, ub, step, [&](Value iv) { - * auto q0 = builder.loadQubit(memref, iv); - * builder.h(q0); - * }); - * ``` - * ```mlir - * scf.for %iv = %lb to %ub step %step { - * %q0 = memref.load %memref[%iv] : memref<3x!qc.qubit> - * qc.h %q0 : !qc.qubit - * } - * ``` - */ + /// Construct an scf.for operation + /// + /// @param lowerbound Lower bound of the loop + /// @param upperbound Upper bound of the loop + /// @param step Step size of the loop + /// @param body Function that builds the body of the for operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfFor(lb, ub, step, [&](Value iv) { + /// auto q0 = builder.loadQubit(memref, iv); + /// builder.h(q0); + /// }); + /// ``` + /// ```mlir + /// scf.for %iv = %lb to %ub step %step { + /// %q0 = memref.load %memref[%iv] : memref<3x!qc.qubit> + /// qc.h %q0 : !qc.qubit + /// } + /// ``` QCProgramBuilder& scfFor(const std::variant& lowerbound, const std::variant& upperbound, const std::variant& step, const function_ref& body); - /** - * @brief Construct an scf.while operation - * - * @param beforeBody Function that builds the before body of the while - * operation - * @param afterBody Function that builds the after body of the while operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfWhile([&] { - * auto res = builder.measure(q0); - * builder.scfCondition(res); - * }, [&] { - * builder.h(q0); - * }); - * ``` - * ```mlir - * scf.while : () -> () { - * %res = qc.measure %q0 : !qc.qubit -> i1 - * scf.condition(%res) - * } do { - * qc.h %q0 : !qc.qubit - * scf.yield - * } - * ``` - */ + /// Construct an scf.while operation + /// + /// @param beforeBody Function that builds the before body of the while + /// operation + /// @param afterBody Function that builds the after body of the while + /// operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfWhile([&] { + /// auto res = builder.measure(q0); + /// builder.scfCondition(res); + /// }, [&] { + /// builder.h(q0); + /// }); + /// ``` + /// ```mlir + /// scf.while : () -> () { + /// %res = qc.measure %q0 : !qc.qubit -> i1 + /// scf.condition(%res) + /// } do { + /// qc.h %q0 : !qc.qubit + /// scf.yield + /// } + /// ``` QCProgramBuilder& scfWhile(const function_ref& beforeBody, const function_ref& afterBody); - /** - * @brief Construct an scf.if operation - * - * @param condition Condition for the if operation - * @param thenBody Function that builds the then body of the if operation - * @param elseBody Function that builds the else body of the if operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfIf(condition, [&] { - * builder.x(q0); - * }, [&] { - * builder.z(q0); - * }); - * ``` - * ```mlir - * scf.if %condition { - * qc.x %q0 : !qc.qubit - * } else { - * qc.z %q0 : !qc.qubit - * } - * ``` - */ + /// Construct an scf.if operation + /// + /// @param condition Condition for the if operation + /// @param thenBody Function that builds the then body of the if operation + /// @param elseBody Function that builds the else body of the if operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfIf(condition, [&] { + /// builder.x(q0); + /// }, [&] { + /// builder.z(q0); + /// }); + /// ``` + /// ```mlir + /// scf.if %condition { + /// qc.x %q0 : !qc.qubit + /// } else { + /// qc.z %q0 : !qc.qubit + /// } + /// ``` QCProgramBuilder& scfIf(const std::variant& condition, const function_ref& thenBody, const function_ref& elseBody = nullptr); - /** - * @brief Construct an scf.if operation conditioned on a classical bit - * - * @details Loads the classical bit from the given classical register at the - * given index and uses it as the condition of the if operation. - * - * @param reg The memref representing the classical register - * @param index The index within the register to load the condition from - * @param thenBody Function that builds the then body of the if operation - * @param elseBody Function that builds the else body of the if operation - * @return Reference to this builder for method chaining - */ + /// Construct an scf.if operation conditioned on a classical bit + /// + /// Loads the classical bit from the given classical register at the + /// given index and uses it as the condition of the if operation. + /// + /// @param reg The memref representing the classical register + /// @param index The index within the register to load the condition from + /// @param thenBody Function that builds the then body of the if operation + /// @param elseBody Function that builds the else body of the if operation + /// @return Reference to this builder for method chaining QCProgramBuilder& scfIf(Value reg, const std::variant& index, const function_ref& thenBody, const function_ref& elseBody = nullptr); - /** - * @brief Construct an scf.index_switch operation - * - * @param arg Index argument. - * @param cases The individual switch cases. - * @param caseBodies An array of functions that build the case bodies. - * @param defaultBody Function that builds the default body. - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfIndexSwitch(index, - * SmallVector{0}, - * SmallVector>{[&] { b.x(q0); }}, - * [&] { b.z(q0); }); - * ``` - * ```mlir - * scf.index_switch %condition - * case 0 { - * qc.x %q0 : !qc.qubit - * } - * default { - * qc.z %q0 : !qc.qubit - * } - * ``` - */ + /// Construct an scf.index_switch operation + /// + /// @param arg Index argument. + /// @param cases The individual switch cases. + /// @param caseBodies An array of functions that build the case bodies. + /// @param defaultBody Function that builds the default body. + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfIndexSwitch(index, + /// SmallVector{0}, + /// SmallVector>{[&] { b.x(q0); }}, + /// [&] { b.z(q0); }); + /// ``` + /// ```mlir + /// scf.index_switch %condition + /// case 0 { + /// qc.x %q0 : !qc.qubit + /// } + /// default { + /// qc.z %q0 : !qc.qubit + /// } + /// ``` QCProgramBuilder& scfIndexSwitch(const std::variant& arg, ArrayRef cases, ArrayRef> caseBodies, const function_ref& defaultBody); - /** - * @brief Construct an scf.condition operation - * - * @param condition Condition for the condition operation - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfCondition(condition); - * ``` - * ```mlir - * scf.condition(%condition) - * ``` - */ + /// Construct an scf.condition operation + /// + /// @param condition Condition for the condition operation + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfCondition(condition); + /// ``` + /// ```mlir + /// scf.condition(%condition) + /// ``` QCProgramBuilder& scfCondition(Value condition); - /** - * @brief Construct an scf.condition operation conditioned on a classical bit - * - * @details Loads the classical bit from the given classical register at the - * given index and uses it as the condition of the condition operation. - * - * @param reg The memref representing the classical register - * @param index The index within the register to load the condition from - * @return Reference to this builder for method chaining - */ + /// Construct an scf.condition operation conditioned on a classical bit + /// + /// Loads the classical bit from the given classical register at the + /// given index and uses it as the condition of the condition operation. + /// + /// @param reg The memref representing the classical register + /// @param index The index within the register to load the condition from + /// @return Reference to this builder for method chaining QCProgramBuilder& scfCondition(Value reg, const std::variant& index); @@ -1329,60 +1281,51 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { // Finalization //===--------------------------------------------------------------------===// - /** - * @brief Finalize the program and return the constructed module - * - * @details - * Automatically deallocates all remaining allocated qubits, adds a return - * statement with exit code 0 (indicating successful execution), and - * transfers ownership of the module to the caller. - * The builder should not be used after calling this method. - * - * @return OwningOpRef containing the constructed quantum program module - */ + /// Finalize the program and return the constructed module + /// + /// Automatically deallocates all remaining allocated qubits, adds a return + /// statement with exit code 0 (indicating successful execution), and + /// transfers ownership of the module to the caller. + /// The builder should not be used after calling this method. + /// + /// @return OwningOpRef containing the constructed quantum program module OwningOpRef finalize(); - /** - * @brief Finalize the program with the given return values and return the - * constructed module - * @param returnValues Values representing the return values of the main - * function. - * - * @details - * Automatically deallocates all remaining valid qubits and tensors of qubits, - * adds a return statement with the given return values, and - * transfers ownership of the module to the caller. The builder should not - * be used after calling this method. - * - * The return values must have the types indicated by the function signature - * of the main function, which returns an `i64` by default and can be - * modified by passing different arguments to the `initialize()` method. - * - * @return OwningOpRef containing the constructed quantum program module - */ + /// Finalize the program with the given return values and return the + /// constructed module + /// @param returnValues Values representing the return values of the main + /// function. + /// + /// Automatically deallocates all remaining valid qubits and tensors of + /// qubits, adds a return statement with the given return values, and + /// transfers ownership of the module to the caller. The builder should not + /// be used after calling this method. + /// + /// The return values must have the types indicated by the function signature + /// of the main function, which returns an `i64` by default and can be + /// modified by passing different arguments to the `initialize()` method. + /// + /// @return OwningOpRef containing the constructed quantum program module OwningOpRef finalize(ValueRange returnValues); - /** - * @brief Convenience method for building quantum programs. - * @param context The MLIR context to use for building the program - * @param buildFunc A function that takes a reference to a QCProgramBuilder - * and uses it to build the desired quantum program. The builder will be - * properly initialized before calling this function, and the resulting module - * will be finalized using the returned Values after this function completes. - * @return The module containing the quantum program built by buildFunc. - */ + /// Convenience method for building quantum programs. + /// @param context The MLIR context to use for building the program + /// @param buildFunc A function that takes a reference to a QCProgramBuilder + /// and uses it to build the desired quantum program. The builder will be + /// properly initialized before calling this function, and the resulting + /// module will be finalized using the returned Values after this function + /// completes. + /// @return The module containing the quantum program built by buildFunc. static OwningOpRef build(MLIRContext* context, const function_ref(QCProgramBuilder&)>& buildFunc); - /** - * @brief Convenience method for building quantum programs with one return - * value. - * @param context The MLIR context to use for building the program - * @param buildFunc A function that takes a reference to a QCProgramBuilder - * and returns the single result value of the desired quantum program. - * @return The module containing the quantum program built by buildFunc. - */ + /// Convenience method for building quantum programs with one return + /// value. + /// @param context The MLIR context to use for building the program + /// @param buildFunc A function that takes a reference to a QCProgramBuilder + /// and returns the single result value of the desired quantum program. + /// @return The module containing the quantum program built by buildFunc. static OwningOpRef build(MLIRContext* context, const function_ref& buildFunc); @@ -1391,7 +1334,7 @@ class QCProgramBuilder final : public ImplicitLocOpBuilder { enum class AllocationMode : uint8_t { Unset, Static, Dynamic }; MLIRContext* ctx{}; - Operation* module; + Operation* moduleOp_; /// Track allocated qubits for automatic deallocation SetVector allocatedQubits; diff --git a/mlir/include/mlir/Dialect/QC/IR/QCOps.h b/mlir/include/mlir/Dialect/QC/IR/QCOps.h index 4ea3810fbf..5daf541ae3 100644 --- a/mlir/include/mlir/Dialect/QC/IR/QCOps.h +++ b/mlir/include/mlir/Dialect/QC/IR/QCOps.h @@ -25,6 +25,9 @@ #include "mlir/Dialect/QC/IR/QCInterfaces.h" #include +#include +#include +#include #include #include diff --git a/mlir/include/mlir/Dialect/QC/IR/QCOps.td b/mlir/include/mlir/Dialect/QC/IR/QCOps.td index e8a1c4cb1a..f900a85873 100644 --- a/mlir/include/mlir/Dialect/QC/IR/QCOps.td +++ b/mlir/include/mlir/Dialect/QC/IR/QCOps.td @@ -16,6 +16,8 @@ include "mlir/Dialect/QC/IR/QCTypes.td" include "mlir/IR/EnumAttr.td" include "mlir/IR/OpBase.td" include "mlir/IR/RegionKindInterface.td" +include "mlir/IR/SymbolInterfaces.td" +include "mlir/Interfaces/CallInterfaces.td" include "mlir/Interfaces/InferTypeOpInterface.td" include "mlir/Interfaces/SideEffectInterfaces.td" @@ -957,6 +959,64 @@ def BarrierOp : QCOp<"barrier", traits = [UnitaryOpInterface]> { }]; } +def CallOp + : QCOp<"call", traits = [CallOpInterface, UnitaryOpInterface, + DeclareOpInterfaceMethods, + MemoryEffects<[MemRead, MemWrite]>]> { + let summary = "Call a unitary QC function"; + let description = [{ + Calls a private `func.func` marked with `mqt.unitary`. The operands match + the callee arguments: zero or more `f64` parameters followed by scalar + qubits. QC reference semantics make the call resultless. + + Example: + ```mlir + qc.call @bell(%q0, %q1) : !qc.qubit, !qc.qubit + ``` + }]; + + let arguments = (ins FlatSymbolRefAttr:$callee, Variadic:$operands, + OptionalAttr:$arg_attrs, + OptionalAttr:$res_attrs); + + let assemblyFormat = [{ + $callee `(` $operands `)` attr-dict `:` type($operands) + }]; + + let builders = [OpBuilder< + (ins "FlatSymbolRefAttr":$callee, "ValueRange":$operands), [{ + $_state.addAttribute("callee", callee); + $_state.addOperands(operands); + }]>]; + + let extraClassDeclaration = [{ + size_t getNumQubits(); + size_t getNumTargets() { return getNumQubits(); } + static size_t getNumControls() { return 0; } + Value getQubit(size_t i) { return getTarget(i); } + Value getTarget(size_t i) { return getQubits()[i]; } + static Value getControl(size_t i) { + llvm::reportFatalUsageError("CallOp does not have controls"); + } + OperandRange getQubits(); + OperandRange getTargets() { return getQubits(); } + static OperandRange getControls() { return {nullptr, 0}; } + size_t getNumParams(); + Value getParameter(size_t i) { return getParameters()[i]; } + OperandRange getParameters(); + StringRef getBaseSymbol() { return getCallee(); } + + ::mlir::Operation::operand_range getArgOperands() { return getOperands(); } + MutableOperandRange getArgOperandsMutable() { return getOperandsMutable(); } + ::mlir::CallInterfaceCallable getCallableForCallee() { + return getCalleeAttr(); + } + void setCalleeFromCallable(::mlir::CallInterfaceCallable callee) { + setCalleeAttr(cast(cast(callee))); + } + }]; +} + //===----------------------------------------------------------------------===// // Modifiers //===----------------------------------------------------------------------===// diff --git a/mlir/include/mlir/Dialect/QCO/Builder/QCOProgramBuilder.h b/mlir/include/mlir/Dialect/QCO/Builder/QCOProgramBuilder.h index 87f1d1a17c..609b8d0f6a 100644 --- a/mlir/include/mlir/Dialect/QCO/Builder/QCOProgramBuilder.h +++ b/mlir/include/mlir/Dialect/QCO/Builder/QCOProgramBuilder.h @@ -14,6 +14,7 @@ #include #include +#include #include #include #include @@ -35,136 +36,138 @@ class ValueRange; namespace qco { -/** - * @brief Builder API for constructing quantum programs in the QCO dialect - * - * @details - * The QCOProgramBuilder provides a type-safe interface for constructing - * quantum circuits using value semantics. Operations consume input qubit - * SSA values and produce new output values, following the functional - * programming paradigm. - * - * @par Linear Type Enforcement: - * The builder enforces linear type semantics by tracking valid qubit SSA - * values. Once a qubit is consumed by an operation producing a new version - * (e.g., reset, measure), the old SSA value is invalidated. This prevents - * use-after-consume errors and mirrors quantum computing's no-cloning theorem. - * - * @par Qubit addressing: - * A program must use either static qubits (`staticQubit`) or dynamic allocation - * (`allocQubit`, `allocQubitRegister`, or `qtensorAlloc`), never both. The - * builder terminates with a usage error if the modes are mixed. - * - * @par Example Usage: - * ```c++ - * QCOProgramBuilder builder(context); - * builder.initialize(); - * - * auto q0 = builder.staticQubit(0); - * auto q1 = builder.staticQubit(1); - * - * // Operations return updated values - * q0 = builder.h(q0); - * std::tie(q0, q1) = builder.cx(q0, q1); - * - * auto module = builder.finalize(); - * ``` - */ +/// Builder API for constructing quantum programs in the QCO dialect +/// +/// The QCOProgramBuilder provides a type-safe interface for constructing +/// quantum circuits using value semantics. Operations consume input qubit +/// SSA values and produce new output values, following the functional +/// programming paradigm. +/// +/// @par Linear Type Enforcement: +/// The builder enforces linear type semantics by tracking valid qubit SSA +/// values. Once a qubit is consumed by an operation producing a new version +/// (e.g., reset, measure), the old SSA value is invalidated. This prevents +/// use-after-consume errors and mirrors quantum computing's no-cloning theorem. +/// +/// @par Qubit addressing: +/// A program must use either static qubits (`staticQubit`) or dynamic +/// allocation +/// (`allocQubit`, `allocQubitRegister`, or `qtensorAlloc`), never both. The +/// builder terminates with a usage error if the modes are mixed. +/// +/// @par Example Usage: +/// ```c++ +/// QCOProgramBuilder builder(context); +/// builder.initialize(); +/// +/// auto q0 = builder.staticQubit(0); +/// auto q1 = builder.staticQubit(1); +/// +/// // Operations return updated values +/// q0 = builder.h(q0); +/// std::tie(q0, q1) = builder.cx(q0, q1); +/// +/// auto moduleOp = builder.finalize(); +/// ``` class QCOProgramBuilder final : public ImplicitLocOpBuilder { public: - /** - * @brief Construct a new QCOProgramBuilder - * @param context The MLIR context to use for building operations - */ + /// Construct a new QCOProgramBuilder + /// @param context The MLIR context to use for building operations explicit QCOProgramBuilder(MLIRContext* context); //===--------------------------------------------------------------------===// // Initialization //===--------------------------------------------------------------------===// - /** - * @brief Initialize the builder and prepare for program construction, with - * a default return type of i64. - * - * @details - * Creates a main function with an `mqt.entry_point` attribute. Must be called - * before adding operations. - */ + /// Initialize the builder and prepare for program construction, with + /// a default return type of i64. + /// + /// Creates a main function with an `mqt.entry_point` attribute. Must be + /// called before adding operations. void initialize(); - /** - * @brief Initialize the builder and prepare for program construction - * with specified return types. - * @param returnTypes The return types for the main function - * - * @details - * Creates a main function with an `mqt.entry_point` attribute. Must be called - * before adding operations. - */ + /// Initialize the builder and prepare for program construction + /// with specified return types. + /// @param returnTypes The return types for the main function + /// + /// Creates a main function with an `mqt.entry_point` attribute. Must be + /// called before adding operations. void initialize(TypeRange returnTypes); - /** - * @brief Modify the return types of the main function after initialization. - * @param returnTypes The new return types for the main function - */ + /// Modify the return types of the main function after initialization. + /// @param returnTypes The new return types for the main function void retype(TypeRange returnTypes); + //===--------------------------------------------------------------------===// + // Functions + //===--------------------------------------------------------------------===// + + /// Create a private function. + /// + /// The callback must return one trailing qubit for every qubit argument, in + /// qubit-argument order. + func::FuncOp + createFunction(StringRef name, TypeRange argumentTypes, + function_ref(ValueRange)> body); + + /// Create a complete private unitary function. + func::FuncOp + createUnitaryFunction(StringRef name, TypeRange argumentTypes, + function_ref(ValueRange)> body); + + /// Call a function, using `qco.call` for a unitary function. + /// + /// Ordinary results are followed by the updated qubit arguments. + SmallVector call(func::FuncOp callee, ValueRange operands); + //===--------------------------------------------------------------------===// // Constants //===--------------------------------------------------------------------===// - /** - * @brief Create a constant integer value - * @param value The value to store in the constant - * @return The value produced by the constant operation - * - * @par Example: - * ```c++ - * auto c = builder.intConstant(1); - * ``` - * ```mlir - * %c = arith.constant 1 : i64 - * ``` - */ + /// Create a constant integer value + /// @param value The value to store in the constant + /// @return The value produced by the constant operation + /// + /// @par Example: + /// ```c++ + /// auto c = builder.intConstant(1); + /// ``` + /// ```mlir + /// %c = arith.constant 1 : i64 + /// ``` Value intConstant(int64_t value); - /** - * @brief Create a constant float value - * @param value The value to store in the constant - * @return The value produced by the constant operation - * - * @par Example: - * ```c++ - * auto c = builder.floatConstant(0.123); - * ``` - * ```mlir - * %c = arith.constant 0.123 : f64 - * ``` - */ + /// Create a constant float value + /// @param value The value to store in the constant + /// @return The value produced by the constant operation + /// + /// @par Example: + /// ```c++ + /// auto c = builder.floatConstant(0.123); + /// ``` + /// ```mlir + /// %c = arith.constant 0.123 : f64 + /// ``` Value floatConstant(double value); - /** - * @brief Create a constant boolean value - * @param value The value to store in the constant - * @return The value produced by the constant operation - * - * @par Example: - * ```c++ - * auto c = builder.boolConstant(true); - * ``` - * ```mlir - * %c = arith.constant 1 : i1 - * ``` - */ + /// Create a constant boolean value + /// @param value The value to store in the constant + /// @return The value produced by the constant operation + /// + /// @par Example: + /// ```c++ + /// auto c = builder.boolConstant(true); + /// ``` + /// ```mlir + /// %c = arith.constant 1 : i1 + /// ``` Value boolConstant(bool value); //===--------------------------------------------------------------------===// // Memory Management //===--------------------------------------------------------------------===// - /** - * @brief A tracked qubit value and its register information. - */ + /// A tracked qubit value and its register information. struct Qubit { /// The tracked SSA value Value value; @@ -173,187 +176,151 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { /// Index of the qubit within its register Value regIndex; - /** - * @brief Implicitly construct a tracked qubit from an SSA value. - * @param value The underlying qubit SSA value - * @param regId ID of the register containing the qubit, or `-1` - * @param regIndex Index of the qubit within its register, if applicable - */ + /// Implicitly construct a tracked qubit from an SSA value. + /// @param value The underlying qubit SSA value + /// @param regId ID of the register containing the qubit, or `-1` + /// @param regIndex Index of the qubit within its register, if applicable // NOLINTNEXTLINE(google-explicit-constructor) Qubit(Value value, int64_t regId = -1, Value regIndex = {}) : value(value), regId(regId), regIndex(regIndex) {} - /** - * @brief Implicitly convert this tracked qubit to its underlying SSA value. - * @return The underlying `Value` - */ + /// Implicitly convert this tracked qubit to its underlying SSA value. + /// @return The underlying `Value` // NOLINTNEXTLINE(google-explicit-constructor) operator Value() const { return value; } - /** - * @brief Get the type of the underlying SSA value. - * @return The underlying value's type - */ + /// Get the type of the underlying SSA value. + /// @return The underlying value's type Type getType() const { return value.getType(); } - /** - * @brief Get the operation defining the underlying SSA value. - * @return The defining operation, or `nullptr` if the value has none - */ + /// Get the operation defining the underlying SSA value. + /// @return The defining operation, or `nullptr` if the value has none Operation* getDefiningOp() const { return value.getDefiningOp(); } - /** - * @brief Get the operation defining the underlying SSA value as @p OpTy. - * @tparam OpTy The expected defining operation type - * @return The defining operation as @p OpTy, or a null operation if the - * value has no defining operation or it is not of type @p OpTy - */ + /// Get the operation defining the underlying SSA value as @p OpTy. + /// @tparam OpTy The expected defining operation type + /// @return The defining operation as @p OpTy, or a null operation if the + /// value has no defining operation or it is not of type @p OpTy template OpTy getDefiningOp() const { return value.getDefiningOp(); } }; - /** - * @brief A tracked qubit tensor value and its register information. - */ + /// A tracked qubit tensor value and its register information. struct Tensor { /// The tracked SSA value Value value; /// ID of the register the tensor corresponds to int64_t regId = -1; - /** - * @brief Implicitly construct a tracked tensor from an SSA value. - * @param value The underlying tensor SSA value - * @param regId ID of the corresponding register, or `-1` - */ + /// Implicitly construct a tracked tensor from an SSA value. + /// @param value The underlying tensor SSA value + /// @param regId ID of the corresponding register, or `-1` // NOLINTNEXTLINE(google-explicit-constructor) Tensor(Value value, int64_t regId = -1) : value(value), regId(regId) {} - /** - * @brief Implicitly convert this tracked tensor to its underlying SSA - * value. - * @return The underlying `Value` - */ + /// Implicitly convert this tracked tensor to its underlying SSA + /// value. + /// @return The underlying `Value` // NOLINTNEXTLINE(google-explicit-constructor) operator Value() const { return value; } - /** - * @brief Get the type of the underlying SSA value. - * @return The underlying value's type - */ + /// Get the type of the underlying SSA value. + /// @return The underlying value's type Type getType() const { return value.getType(); } - /** - * @brief Get the operation defining the underlying SSA value. - * @return The defining operation, or `nullptr` if the value has none - */ + /// Get the operation defining the underlying SSA value. + /// @return The defining operation, or `nullptr` if the value has none Operation* getDefiningOp() const { return value.getDefiningOp(); } - /** - * @brief Get the operation defining the underlying SSA value as @p OpTy. - * @tparam OpTy The expected defining operation type - * @return The defining operation as @p OpTy, or a null operation if the - * value has no defining operation or it is not of type @p OpTy - */ + /// Get the operation defining the underlying SSA value as @p OpTy. + /// @tparam OpTy The expected defining operation type + /// @return The defining operation as @p OpTy, or a null operation if the + /// value has no defining operation or it is not of type @p OpTy template OpTy getDefiningOp() const { return value.getDefiningOp(); } }; - /** - * @brief Represents a qubit register with its qubits. - */ + /// Represents a qubit register with its qubits. struct QubitRegister { /// The QTensor value representing the qubit register Value value; /// The allocated qubit values SmallVector qubits; - /** - * @brief Access a specific qubit in the register - * @param index The index of the qubit to access - * @return The specified qubit value - */ + /// Access a specific qubit in the register + /// @param index The index of the qubit to access + /// @return The specified qubit value Value& operator[](size_t index); - /** - * @brief Conversion to the backing QTensor value - * @return The QTensor value representing the qubit register - */ + /// Conversion to the backing QTensor value + /// @return The QTensor value representing the qubit register explicit operator Value() const { return value; } }; - /** - * @brief Allocate a single qubit initialized to |0⟩ - * @return A tracked qubit handle (convertible to `Value`) - * - * @par Example: - * ```c++ - * auto q = builder.allocQubit(); - * ``` - * ```mlir - * %q = qco.alloc : !qco.qubit - * ``` - */ + /// Allocate a single qubit initialized to |0⟩ + /// @return A tracked qubit handle (convertible to `Value`) + /// + /// @par Example: + /// ```c++ + /// auto q = builder.allocQubit(); + /// ``` + /// ```mlir + /// %q = qco.alloc : !qco.qubit + /// ``` Qubit allocQubit(); - /** - * @brief Get a static qubit by index - * @param index The qubit index - * @return A tracked qubit handle (convertible to `Value`) - * - * @par Example: - * ```c++ - * auto q0 = builder.staticQubit(0); - * ``` - * ```mlir - * %q0 = qco.static 0 : !qco.qubit - * ``` - */ + /// Get a static qubit by index + /// @param index The qubit index + /// @return A tracked qubit handle (convertible to `Value`) + /// + /// @par Example: + /// ```c++ + /// auto q0 = builder.staticQubit(0); + /// ``` + /// ```mlir + /// %q0 = qco.static 0 : !qco.qubit + /// ``` Qubit staticQubit(uint64_t index); - /** - * @brief Allocate a qubit tensor and eagerly extract every element - * @param size Number of qubits (must be positive) - * @param name Optional source-level register name - * @return A `QubitRegister` containing the residual tensor and one standalone - * qubit value for every eagerly extracted element - * - * @par Example: - * ```c++ - * auto q = builder.allocQubitRegister(3); - * ``` - * ```mlir - * %t0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - * %t1, %q0 = qtensor.extract %t0[%c0]: tensor<3x!qco.qubit> - * %t2, %q1 = qtensor.extract %t1[%c1]: tensor<3x!qco.qubit> - * %t3, %q2 = qtensor.extract %t2[%c2]: tensor<3x!qco.qubit> - * ``` - */ + /// Allocate a qubit tensor and eagerly extract every element + /// @param size Number of qubits (must be positive) + /// @param name Optional source-level register name + /// @return A `QubitRegister` containing the residual tensor and one + /// standalone qubit value for every eagerly extracted element + /// + /// @par Example: + /// ```c++ + /// auto q = builder.allocQubitRegister(3); + /// ``` + /// ```mlir + /// %t0 = qtensor.alloc(%c3) : tensor<3x!qco.qubit> + /// %t1, %q0 = qtensor.extract %t0[%c0]: tensor<3x!qco.qubit> + /// %t2, %q1 = qtensor.extract %t1[%c1]: tensor<3x!qco.qubit> + /// %t3, %q2 = qtensor.extract %t2[%c2]: tensor<3x!qco.qubit> + /// ``` QubitRegister allocQubitRegister(int64_t size, StringRef name = {}); - /** - * @brief Allocate a classical bit register - * - * @details The register uses `!cbit.reg`. Its initialization is explicit - * and independent of every other register built by this builder. - * - * @param size Number of bits (must be positive) - * @param name Optional source-level register name; defaults to no name - * @param initialization Initial value of the register elements; defaults to - * zero - * @return The CBit register value - * - * @par Example: - * ```c++ - * auto c = builder.allocClassicalBitRegister(3, "c"); - * ``` - * ```mlir - * %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} - * : !cbit.reg<3> - * ``` - */ + /// Allocate a classical bit register + /// + /// The register uses `!cbit.reg`. Its initialization is explicit + /// and independent of every other register built by this builder. + /// + /// @param size Number of bits (must be positive) + /// @param name Optional source-level register name; defaults to no name + /// @param initialization Initial value of the register elements; defaults to + /// zero + /// @return The CBit register value + /// + /// @par Example: + /// ```c++ + /// auto c = builder.allocClassicalBitRegister(3, "c"); + /// ``` + /// ```mlir + /// %c = cbit.alloc(#cbit.init) {mqt.register_name = "c"} + /// : !cbit.reg<3> + /// ``` Value allocClassicalBitRegister( int64_t size, StringRef name = {}, cbit::Initialization initialization = cbit::Initialization::Zero); @@ -369,187 +336,164 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { // QTensor operations //===--------------------------------------------------------------------===// - /** - * @brief Allocate a qubit tensor - * - * @details Allocates and returns one intact, one-dimensional tensor of - * `!qco.qubit` values. No elements are extracted. If the size is a constant, - * the tensor has static size; otherwise it has dynamic size. Its qubits are - * initialized in the |0> state, and the tensor is tracked automatically. - * - * @param size Number of qubits (must be positive) - * @return The allocated tensor - * - * @par Example: - * ```c++ - * auto tensor = builder.qtensorAlloc(3); - * ``` - * ```mlir - * %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - * ``` - */ + /// Allocate a qubit tensor + /// + /// Allocates and returns one intact, one-dimensional tensor of + /// `!qco.qubit` values. No elements are extracted. If the size is a constant, + /// the tensor has static size; otherwise it has dynamic size. Its qubits are + /// initialized in the |0> state, and the tensor is tracked automatically. + /// + /// @param size Number of qubits (must be positive) + /// @return The allocated tensor + /// + /// @par Example: + /// ```c++ + /// auto tensor = builder.qtensorAlloc(3); + /// ``` + /// ```mlir + /// %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> + /// ``` Value qtensorAlloc(const std::variant& size); - /** - * @brief Allocate a qubit tensor from a list of qubit values - * - * @details - * Consumes the input qubits and creates a one-dimensional tensor of - * !qco.qubit types. The resulting tensor has a static size given by the - * number of input values. The consumed qubits are removed from the qubit - * tracking and the resulting tensor is added to the tracking. - * - * @param elements Inserted Qubits (must be valid/unconsumed) - * @return The allocated tensor - * - * @par Example: - * ```c++ - * auto tensor = builder.qtensorFromElements({q0, q1, q2}); - * ``` - * ```mlir - * %tensor = qtensor.from_elements %q0, %q1, %q2 : tensor<3x!qco.qubit> - * ``` - */ + /// Allocate a qubit tensor from a list of qubit values + /// + /// Consumes the input qubits and creates a one-dimensional tensor of + /// !qco.qubit types. The resulting tensor has a static size given by the + /// number of input values. The consumed qubits are removed from the qubit + /// tracking and the resulting tensor is added to the tracking. + /// + /// @param elements Inserted Qubits (must be valid/unconsumed) + /// @return The allocated tensor + /// + /// @par Example: + /// ```c++ + /// auto tensor = builder.qtensorFromElements({q0, q1, q2}); + /// ``` + /// ```mlir + /// %tensor = qtensor.from_elements %q0, %q1, %q2 : tensor<3x!qco.qubit> + /// ``` Value qtensorFromElements(ValueRange elements); - /** - * @brief Extract a qubit from a tensor - * - * @details - * Extracts a qubit from a one-dimensional tensor of qubits at the given index - * and returns the updated tensor and the extracted qubit. The extracted qubit - * is added to the qubit tracking and the tracking of the source tensor is - * updated. - * - * @param tensor Source tensor (must be valid/unconsumed) - * @param index The index from where the qubit is extracted - * @return Pair of (outTensor, extractedQubit) - * - * @par Example: - * ```c++ - * auto [outTensor, q0] = builder.qtensorExtract(tensor, 0); - * ``` - * ```mlir - * %outTensor, %q0 = qtensor.extract %tensor[%c0]: tensor<3x!qco.qubit> - * ``` - */ + /// Extract a qubit from a tensor + /// + /// Extracts a qubit from a one-dimensional tensor of qubits at the given + /// index and returns the updated tensor and the extracted qubit. The + /// extracted qubit is added to the qubit tracking and the tracking of the + /// source tensor is updated. + /// + /// @param tensor Source tensor (must be valid/unconsumed) + /// @param index The index from where the qubit is extracted + /// @return Pair of (outTensor, extractedQubit) + /// + /// @par Example: + /// ```c++ + /// auto [outTensor, q0] = builder.qtensorExtract(tensor, 0); + /// ``` + /// ```mlir + /// %outTensor, %q0 = qtensor.extract %tensor[%c0]: tensor<3x!qco.qubit> + /// ``` std::pair qtensorExtract(Value tensor, const std::variant& index); - /** - * @brief Insert a qubit into a tensor - * - * @details - * Inserts a scalar qubit into the one-dimensional tensor of qubits at the - * given index. The inserted qubit is consumed and removed from the qubit - * tracking while the tracking for the source tensor is updated. - * - * @param scalar The scalar qubit that is inserted (must be valid/unconsumed) - * @param tensor The tensor where the qubit is inserted (must be - * valid/unconsumed) - * @param index The index into where the qubit is inserted - * @return The output tensor - * - * @par Example: - * ```c++ - * auto outTensor = builder.qtensorInsert(q0, tensor, 0); - * ``` - * ```mlir - * %outTensor = qtensor.insert %q0 into %tensor[%c0] : tensor<3x!qco.qubit> - * ``` - */ + /// Insert a qubit into a tensor + /// + /// Inserts a scalar qubit into the one-dimensional tensor of qubits at the + /// given index. The inserted qubit is consumed and removed from the qubit + /// tracking while the tracking for the source tensor is updated. + /// + /// @param scalar The scalar qubit that is inserted (must be valid/unconsumed) + /// @param tensor The tensor where the qubit is inserted (must be + /// valid/unconsumed) + /// @param index The index into where the qubit is inserted + /// @return The output tensor + /// + /// @par Example: + /// ```c++ + /// auto outTensor = builder.qtensorInsert(q0, tensor, 0); + /// ``` + /// ```mlir + /// %outTensor = qtensor.insert %q0 into %tensor[%c0] : tensor<3x!qco.qubit> + /// ``` Value qtensorInsert(Value scalar, Value tensor, const std::variant& index); - /** - * @brief Explicitly deallocate a tensor - * - * @details - * Validates and removes the tensor from tracking. Qubits or tensors of qubits - * that were extracted from the tensor but not inserted back again need to be - * deallocated separately. Optional; `finalize()` automatically deallocates - * all remaining tensors. - * - * @param tensor Tensor to deallocate (must be valid/unconsumed) - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.qtensorDealloc(tensor); - * ``` - * ```mlir - * qtensor.dealloc %tensor : tensor<3x!qco.qubit> - * ``` - */ + /// Explicitly deallocate a tensor + /// + /// Validates and removes the tensor from tracking. Qubits or tensors of + /// qubits that were extracted from the tensor but not inserted back again + /// need to be deallocated separately. Optional; `finalize()` automatically + /// deallocates all remaining tensors. + /// + /// @param tensor Tensor to deallocate (must be valid/unconsumed) + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.qtensorDealloc(tensor); + /// ``` + /// ```mlir + /// qtensor.dealloc %tensor : tensor<3x!qco.qubit> + /// ``` QCOProgramBuilder& qtensorDealloc(Value tensor); //===--------------------------------------------------------------------===// // Measurement and Reset //===--------------------------------------------------------------------===// - /** - * @brief Measure a qubit in the computational basis - * - * @details - * Consumes the input qubit and produces a new output qubit SSA value - * along with the measurement result (i1). The input is validated and - * tracking is updated to reflect the new output value. - * - * @param qubit Input qubit (must be valid/unconsumed) - * @return Pair of (output_qubit, measurement_result) - * - * @par Example: - * ```c++ - * auto [q_out, result] = builder.measure(q); - * ``` - * ```mlir - * %q_out, %result = qco.measure %q : !qco.qubit - * ``` - */ + /// Measure a qubit in the computational basis + /// + /// Consumes the input qubit and produces a new output qubit SSA value + /// along with the measurement result (i1). The input is validated and + /// tracking is updated to reflect the new output value. + /// + /// @param qubit Input qubit (must be valid/unconsumed) + /// @return Pair of (output_qubit, measurement_result) + /// + /// @par Example: + /// ```c++ + /// auto [q_out, result] = builder.measure(q); + /// ``` + /// ```mlir + /// %q_out, %result = qco.measure %q : !qco.qubit + /// ``` std::pair measure(Value qubit); - /** - * @brief Measure a qubit and store the result in a classical bit register - * - * @details - * Measures the qubit and stores the classical result in the given classical - * register at the given index, in addition to returning it. - * - * @param qubit Input qubit (must be valid/unconsumed) - * @param reg The CBit register - * @param index The index within the classical register - * @return Pair of (output_qubit, measurement_result) - * - * @par Example: - * ```c++ - * auto [q0_out, r0] = builder.measure(q0, c, 0); - * ``` - * ```mlir - * %q0_out, %r0 = qco.measure %q0 : !qco.qubit - * cbit.store %r0, %c[%c0] : !cbit.reg<3> - * ``` - */ + /// Measure a qubit and store the result in a classical bit register + /// + /// Measures the qubit and stores the classical result in the given classical + /// register at the given index, in addition to returning it. + /// + /// @param qubit Input qubit (must be valid/unconsumed) + /// @param reg The CBit register + /// @param index The index within the classical register + /// @return Pair of (output_qubit, measurement_result) + /// + /// @par Example: + /// ```c++ + /// auto [q0_out, r0] = builder.measure(q0, c, 0); + /// ``` + /// ```mlir + /// %q0_out, %r0 = qco.measure %q0 : !qco.qubit + /// cbit.store %r0, %c[%c0] : !cbit.reg<3> + /// ``` std::pair measure(Value qubit, Value reg, const std::variant& index); - /** - * @brief Reset a qubit to |0⟩ state - * - * @details - * Consumes the input qubit and produces a new output qubit SSA value - * in the |0⟩ state. The input is validated and tracking is updated. - * - * @param qubit Input qubit (must be valid/unconsumed) - * @return Output qubit value - * - * @par Example: - * ```c++ - * q = builder.reset(q); - * ``` - * ```mlir - * %q_out = qco.reset %q : !qco.qubit -> !qco.qubit - * ``` - */ + /// Reset a qubit to |0⟩ state + /// + /// Consumes the input qubit and produces a new output qubit SSA value + /// in the |0⟩ state. The input is validated and tracking is updated. + /// + /// @param qubit Input qubit (must be valid/unconsumed) + /// @return Output qubit value + /// + /// @par Example: + /// ```c++ + /// q = builder.reset(q); + /// ``` + /// ```mlir + /// %q_out = qco.reset %q : !qco.qubit -> !qco.qubit + /// ``` Value reset(Value qubit); //===--------------------------------------------------------------------===// @@ -1368,188 +1312,170 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { // BarrierOp - /** - * @brief Apply a BarrierOp - * - * @param qubits Input qubits (must be valid/unconsumed) - * @return Output qubits - * - * @par Example: - * ```c++ - * builder.barrier({q0, q1}); - * ``` - * ```mlir - * qco.barrier %q0, %q1 : !qco.qubit, !qco.qubit -> !qco.qubit, - * !qco.qubit - * ``` - */ + /// Apply a BarrierOp + /// + /// @param qubits Input qubits (must be valid/unconsumed) + /// @return Output qubits + /// + /// @par Example: + /// ```c++ + /// builder.barrier({q0, q1}); + /// ``` + /// ```mlir + /// qco.barrier %q0, %q1 : !qco.qubit, !qco.qubit -> !qco.qubit, + /// !qco.qubit + /// ``` ValueRange barrier(ValueRange qubits); - /** - * @brief Apply an explicitly represented dense unitary matrix - * - * @param qubits Input qubits (must be valid/unconsumed), ordered from the - * most-significant basis bit to the least-significant basis bit - * @param matrix Square row-major `complex` matrix - * @return Output qubits - */ + /// Apply an explicitly represented dense unitary matrix + /// + /// @param qubits Input qubits (must be valid/unconsumed), ordered from the + /// most-significant basis bit to the least-significant basis bit + /// @param matrix Square row-major `complex` matrix + /// @return Output qubits ValueRange unitary(ValueRange qubits, DenseElementsAttr matrix); //===--------------------------------------------------------------------===// // Modifiers //===--------------------------------------------------------------------===// - /** - * @brief Apply a control modifier to a collection of gates - * - * @param controls Input control qubits - * @param targets Input target qubits - * @param body Function that builds the body containing the target gates - * @return Pair of (output_control_qubits, output_target_qubits) - * - * @par Example: - * ```c++ - * auto [controls_out, targets_out] = - * builder.ctrl(q0_in, q1_in, - * [&](ValueRange targets) -> SmallVector { - * return {builder.x(targets[0])}; - * }); - * ``` - * ```mlir - * %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%t = %q1_in) { - * %q1_res = qco.x %t : !qco.qubit -> !qco.qubit - * qco.yield %q1_res : !qco.qubit - * } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) - * ``` - */ + /// Apply a control modifier to a collection of gates + /// + /// @param controls Input control qubits + /// @param targets Input target qubits + /// @param body Function that builds the body containing the target gates + /// @return Pair of (output_control_qubits, output_target_qubits) + /// + /// @par Example: + /// ```c++ + /// auto [controls_out, targets_out] = + /// builder.ctrl(q0_in, q1_in, + /// [&](ValueRange targets) -> SmallVector { + /// return {builder.x(targets[0])}; + /// }); + /// ``` + /// ```mlir + /// %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%t = %q1_in) { + /// %q1_res = qco.x %t : !qco.qubit -> !qco.qubit + /// qco.yield %q1_res : !qco.qubit + /// } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) + /// ``` std::pair ctrl(ValueRange controls, ValueRange targets, function_ref(ValueRange)> body); - /** - * @brief Apply a control modifier with a single target and one-qubit body. - * - * @param controls Control qubits - * @param target Target qubit - * @param body Function that builds the body containing the target operation - * @return Pair of (output_control_qubits, output_target_qubit) - * - * @par Example: - * ```c++ - * auto [controls_out, target_out] = - * builder.ctrl({q0_in, q1_in}, q2_in, [&](Value target) { - * return builder.x(target); - * }); - * ``` - */ + /// Apply a control modifier with a single target and one-qubit body. + /// + /// @param controls Control qubits + /// @param target Target qubit + /// @param body Function that builds the body containing the target operation + /// @return Pair of (output_control_qubits, output_target_qubit) + /// + /// @par Example: + /// ```c++ + /// auto [controls_out, target_out] = + /// builder.ctrl({q0_in, q1_in}, q2_in, [&](Value target) { + /// return builder.x(target); + /// }); + /// ``` std::pair ctrl(ValueRange controls, Value target, function_ref body); - /** - * @brief Apply a control modifier with one control and one target. - * - * @param control Control qubit - * @param target Target qubit - * @param body Function that builds the body containing the target operation - * @return Pair of (output_control_qubit, output_target_qubit) - * - * @par Example: - * ```c++ - * auto [control_out, target_out] = - * builder.ctrl(q0_in, q1_in, [&](Value target) { - * return builder.x(target); - * }); - * ``` - */ + /// Apply a control modifier with one control and one target. + /// + /// @param control Control qubit + /// @param target Target qubit + /// @param body Function that builds the body containing the target operation + /// @return Pair of (output_control_qubit, output_target_qubit) + /// + /// @par Example: + /// ```c++ + /// auto [control_out, target_out] = + /// builder.ctrl(q0_in, q1_in, [&](Value target) { + /// return builder.x(target); + /// }); + /// ``` std::pair ctrl(Value control, Value target, function_ref body); - /** - * @brief Apply an inverse (i.e., adjoint) modifier to a collection of gates - * - * @param qubits Input qubits - * @param body Function that builds the body containing the gates to invert - * @return Output qubits - * - * @par Example: - * ```c++ - * auto qubits_out = builder.inv(q0_in, - * [&](ValueRange qubits) -> SmallVector { - * return {builder.s(qubits[0])}; - * } - * ); - * ``` - * ```mlir - * %qubits_out = qco.inv (%q = %q0_in) { - * %q_res = qco.s %q : !qco.qubit -> !qco.qubit - * qco.yield %q_res : !qco.qubit - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - */ + /// Apply an inverse (i.e., adjoint) modifier to a collection of gates + /// + /// @param qubits Input qubits + /// @param body Function that builds the body containing the gates to invert + /// @return Output qubits + /// + /// @par Example: + /// ```c++ + /// auto qubits_out = builder.inv(q0_in, + /// [&](ValueRange qubits) -> SmallVector { + /// return {builder.s(qubits[0])}; + /// } + /// ); + /// ``` + /// ```mlir + /// %qubits_out = qco.inv (%q = %q0_in) { + /// %q_res = qco.s %q : !qco.qubit -> !qco.qubit + /// qco.yield %q_res : !qco.qubit + /// } : {!qco.qubit} -> {!qco.qubit} + /// ``` ValueRange inv(ValueRange qubits, function_ref(ValueRange)> body); - /** - * @brief Apply an inverse modifier on a single qubit. - * - * @param qubit Qubit involved in the operation - * @param body Function that builds the body containing the operation to - * invert - * @return Output qubit - * - * @par Example: - * ```c++ - * auto qubit_out = builder.inv(q0_in, [&](Value qubit) { - * return builder.s(qubit); - * }); - * ``` - */ + /// Apply an inverse modifier on a single qubit. + /// + /// @param qubit Qubit involved in the operation + /// @param body Function that builds the body containing the operation to + /// invert + /// @return Output qubit + /// + /// @par Example: + /// ```c++ + /// auto qubit_out = builder.inv(q0_in, [&](Value qubit) { + /// return builder.s(qubit); + /// }); + /// ``` Value inv(Value qubit, function_ref body); - /** - * @brief Apply a power modifier to a collection of gates - * - * @param exponent The exponent to raise the gates to - * @param qubits Input qubits - * @param body Function that builds the body containing the gates to - * exponentiate - * @return Output qubits - * - * @par Example: - * ```c++ - * qubits_out = builder.pow(2.0, {q0_in, q1_in}, - * [&](ValueRange qubits) -> SmallVector { - * auto [q0, q1] = builder.swap(qubits[0], qubits[1]); - * return {q0, q1}; - * } - * ); - * ``` - * ```mlir - * %q_out = qco.pow(%exponent) (%q = %q_in) { - * %q_res = qco.s %q : !qco.qubit -> !qco.qubit - * qco.yield %q_res - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - */ + /// Apply a power modifier to a collection of gates + /// + /// @param exponent The exponent to raise the gates to + /// @param qubits Input qubits + /// @param body Function that builds the body containing the gates to + /// exponentiate + /// @return Output qubits + /// + /// @par Example: + /// ```c++ + /// qubits_out = builder.pow(2.0, {q0_in, q1_in}, + /// [&](ValueRange qubits) -> SmallVector { + /// auto [q0, q1] = builder.swap(qubits[0], qubits[1]); + /// return {q0, q1}; + /// } + /// ); + /// ``` + /// ```mlir + /// %q_out = qco.pow(%exponent) (%q = %q_in) { + /// %q_res = qco.s %q : !qco.qubit -> !qco.qubit + /// qco.yield %q_res + /// } : {!qco.qubit} -> {!qco.qubit} + /// ``` ValueRange pow(const std::variant& exponent, ValueRange qubits, function_ref(ValueRange)> body); - /** - * @brief Apply a power modifier on a single qubit. - * - * @param exponent The exponent to raise the operation to - * @param qubit Input qubit - * @param body Function that builds the body containing the operation to - * exponentiate - * @return Output qubit - * - * @par Example: - * ```c++ - * auto qubit_out = builder.pow(2.0, q0_in, [&](Value qubit) { - * return builder.s(qubit); - * }); - * ``` - */ + /// Apply a power modifier on a single qubit. + /// + /// @param exponent The exponent to raise the operation to + /// @param qubit Input qubit + /// @param body Function that builds the body containing the operation to + /// exponentiate + /// @return Output qubit + /// + /// @par Example: + /// ```c++ + /// auto qubit_out = builder.pow(2.0, q0_in, [&](Value qubit) { + /// return builder.s(qubit); + /// }); + /// ``` Value pow(const std::variant& exponent, Value qubit, function_ref body); @@ -1557,327 +1483,302 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { // Deallocation //===--------------------------------------------------------------------===// - /** - * @brief Consume a qubit value (end of lifetime) - * - * @details - * Validates and removes the qubit from tracking. Optional; `finalize()` - * automatically sinks all remaining qubits. - * - * @param qubit Qubit to sink (must be valid/unconsumed) - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.sink(q); - * ``` - * ```mlir - * qco.sink %q : !qco.qubit - * ``` - */ + /// Consume a qubit value (end of lifetime) + /// + /// Validates and removes the qubit from tracking. Optional; `finalize()` + /// automatically sinks all remaining qubits. + /// + /// @param qubit Qubit to sink (must be valid/unconsumed) + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.sink(q); + /// ``` + /// ```mlir + /// qco.sink %q : !qco.qubit + /// ``` QCOProgramBuilder& sink(Value qubit); //===--------------------------------------------------------------------===// // SCF operations //===--------------------------------------------------------------------===// - /** - * @brief Construct an if operation for qubits or tensors of qubits with - * linear typing - * - * @details - * Constructs an if operation that takes a bool Value and a range of qubit - * and qtensor values that are used in the then/else region of this operation. - * The values are passed down as block arguments to each region. Qubits that - * were extracted from a tensor that is used as an argument for this operation - * are automatically inserted before the operation is constructed. - * - * @param condition Bool condition - * @param initArgs Initial arguments for the if branches - * @param thenBody Function that builds the then body of the if operation - * @param elseBody Function that builds the else body of the if operation - * @return ValueRange of the results - * - * @par Example: - * ```c++ - * result = builder.qcoIf(condition, initArgs, [&](ValueRange args) - * -> SmallVector { - * auto q1 = builder.x(args[0]); - * return {q1}; - * }, [&](ValueRange args) -> SmallVector { - * auto q2 = builder.z(args[0]); - * return {q2}; - * }); - * ``` - * ```mlir - * %q3 = qco.if %condition args(%arg0 = %q0) -> (!qco.qubit) { - * %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } else args(%arg0 = %q0) { - * %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q2 : !qco.qubit - * } - * ``` - */ + /// Construct an if operation for qubits or tensors of qubits with + /// linear typing + /// + /// Constructs an if operation that takes a bool Value and a range of qubit + /// and qtensor values that are used in the then/else region of this + /// operation. The values are passed down as block arguments to each region. + /// Qubits that were extracted from a tensor that is used as an argument for + /// this operation are automatically inserted before the operation is + /// constructed. + /// + /// @param condition Bool condition + /// @param initArgs Initial arguments for the if branches + /// @param thenBody Function that builds the then body of the if operation + /// @param elseBody Function that builds the else body of the if operation + /// @return ValueRange of the results + /// + /// @par Example: + /// ```c++ + /// result = builder.qcoIf(condition, initArgs, [&](ValueRange args) + /// -> SmallVector { + /// auto q1 = builder.x(args[0]); + /// return {q1}; + /// }, [&](ValueRange args) -> SmallVector { + /// auto q2 = builder.z(args[0]); + /// return {q2}; + /// }); + /// ``` + /// ```mlir + /// %q3 = qco.if %condition args(%arg0 = %q0) -> (!qco.qubit) { + /// %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q1 : !qco.qubit + /// } else args(%arg0 = %q0) { + /// %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q2 : !qco.qubit + /// } + /// ``` ValueRange qcoIf(const std::variant& condition, ValueRange initArgs, function_ref(ValueRange)> thenBody, function_ref(ValueRange)> elseBody = nullptr); - /** - * @brief Construct an scf.if operation conditioned on a classical bit - * - * @details Loads the classical bit from the given classical register at the - * given index and uses it as the condition of the if operation. - * - * @param reg The CBit register - * @param index The index within the register to load the condition from - * @param initArgs Initial arguments threaded through the if operation - * @param thenBody Function that builds the then body of the if operation - * @param elseBody Function that builds the else body of the if operation - * @return ValueRange of the results - */ + /// Construct an scf.if operation conditioned on a classical bit + /// + /// Loads the classical bit from the given classical register at the + /// given index and uses it as the condition of the if operation. + /// + /// @param reg The CBit register + /// @param index The index within the register to load the condition from + /// @param initArgs Initial arguments threaded through the if operation + /// @param thenBody Function that builds the then body of the if operation + /// @param elseBody Function that builds the else body of the if operation + /// @return ValueRange of the results ValueRange qcoIf(Value reg, const std::variant& index, ValueRange initArgs, function_ref(ValueRange)> thenBody, function_ref(ValueRange)> elseBody = nullptr); - /** - * @brief Construct an if operation for qubits with a single target qubit or - * tensor. - * - * @details - * Constructs an if operation that takes a bool Value and a single qubit - * or qtensor value that is used in the then/else region of this operation. - * The value is passed down as block arguments to each region. Qubits that - * were extracted from a tensor that is used as an argument for this operation - * are automatically inserted before the operation is constructed. - * - * @param condition Bool condition - * @param initArg Initial argument for the if branches - * @param thenBody Function that builds the then body of the if operation - * @param elseBody Function that builds the else body of the if operation - * @return Value as a result - * - * @par Example: - * ```c++ - * result = builder.qcoIf(condition, initArg, [&](Value arg) - * -> Value { - * auto q1 = builder.x(arg); - * return q1; - * }, [&](Value arg) -> Value { - * auto q2 = builder.z(arg); - * return q2; - * }); - * ``` - * ```mlir - * %q3 = qco.if %condition args(%arg0 = %q0) -> (!qco.qubit) { - * %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } else args(%arg0 = %q0) { - * %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q2 : !qco.qubit - * } - * ``` - */ + /// Construct an if operation for qubits with a single target qubit or + /// tensor. + /// + /// Constructs an if operation that takes a bool Value and a single qubit + /// or qtensor value that is used in the then/else region of this operation. + /// The value is passed down as block arguments to each region. Qubits that + /// were extracted from a tensor that is used as an argument for this + /// operation are automatically inserted before the operation is constructed. + /// + /// @param condition Bool condition + /// @param initArg Initial argument for the if branches + /// @param thenBody Function that builds the then body of the if operation + /// @param elseBody Function that builds the else body of the if operation + /// @return Value as a result + /// + /// @par Example: + /// ```c++ + /// result = builder.qcoIf(condition, initArg, [&](Value arg) + /// -> Value { + /// auto q1 = builder.x(arg); + /// return q1; + /// }, [&](Value arg) -> Value { + /// auto q2 = builder.z(arg); + /// return q2; + /// }); + /// ``` + /// ```mlir + /// %q3 = qco.if %condition args(%arg0 = %q0) -> (!qco.qubit) { + /// %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q1 : !qco.qubit + /// } else args(%arg0 = %q0) { + /// %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q2 : !qco.qubit + /// } + /// ``` Value qcoIf(const std::variant& condition, Value initArg, function_ref thenBody, function_ref elseBody = nullptr); - /** - * @brief Construct an index switch operation for qubits or tensors of qubits - * with linear typing. - * - * @details - * Constructs an index switch operation that takes an index Value and a range - * of qubit and qtensor values that are used in the case regions of this - * operation. The values are passed down as block arguments to each region. - * Qubits that were extracted from a tensor that is used as an argument for - * this operation are automatically inserted before the operation is - * constructed. - * - * @param arg Index argument. - * @param targets Initial arguments for the index switch branches. - * @param cases The individual switch cases. - * @param caseBodies An array of functions that build the case bodies. - * @param defaultBody Function that builds the default body. - * @return ValueRange of the results. - * - * @par Example: - * ```c++ - * result = b.qcoIndexSwitch(arg, initTargets, - * SmallVector{0}, - * SmallVector(ValueRange)>>{ - * [&](ValueRange args) { - * auto q1 = builder.x(args[0]); - * return {q1}; - * } - * }, - * [&](ValueRange args) { - * auto q2 = builder.x(args[0]); - * return {q2}; - * }); - * ``` - * ```mlir - * %result = qco.index_switch %arg -> !qco.qubit - * case 0 args(%arg0 = %q0) { - * %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } - * default args(%arg0 = %q0) { - * %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q2 : !qco.qubit - * } - * ``` - */ + /// Construct an index switch operation for qubits or tensors of qubits + /// with linear typing. + /// + /// Constructs an index switch operation that takes an index Value and a range + /// of qubit and qtensor values that are used in the case regions of this + /// operation. The values are passed down as block arguments to each region. + /// Qubits that were extracted from a tensor that is used as an argument for + /// this operation are automatically inserted before the operation is + /// constructed. + /// + /// @param arg Index argument. + /// @param targets Initial arguments for the index switch branches. + /// @param cases The individual switch cases. + /// @param caseBodies An array of functions that build the case bodies. + /// @param defaultBody Function that builds the default body. + /// @return ValueRange of the results. + /// + /// @par Example: + /// ```c++ + /// result = b.qcoIndexSwitch(arg, initTargets, + /// SmallVector{0}, + /// SmallVector(ValueRange)>>{ + /// [&](ValueRange args) { + /// auto q1 = builder.x(args[0]); + /// return {q1}; + /// } + /// }, + /// [&](ValueRange args) { + /// auto q2 = builder.x(args[0]); + /// return {q2}; + /// }); + /// ``` + /// ```mlir + /// %result = qco.index_switch %arg -> !qco.qubit + /// case 0 args(%arg0 = %q0) { + /// %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q1 : !qco.qubit + /// } + /// default args(%arg0 = %q0) { + /// %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit + /// qco.yield %q2 : !qco.qubit + /// } + /// ``` ValueRange qcoIndexSwitch( const std::variant& arg, ValueRange targets, ArrayRef cases, ArrayRef(ValueRange)>> caseBodies, function_ref(ValueRange)> defaultBody); - /** - * @brief Construct an index switch operation with a single linear target. - * - * @details - * Constructs an index switch operation for one qubit or qtensor value. - * Each branch callback receives and returns a single value, avoiding - * one-element ranges and vectors. - * - * @param arg Index argument. - * @param target Initial argument for every index switch branch. - * @param cases The individual switch cases. - * @param caseBodies Functions that build the case bodies. - * @param defaultBody Function that builds the default body. - * @return The single result value. - * - * @par Example: - * ```c++ - * result = builder.qcoIndexSwitch( - * arg, target, SmallVector{0}, - * SmallVector>{ - * [&](Value value) { return builder.x(value); }}, - * [&](Value value) { return builder.z(value); }); - * ``` - */ + /// Construct an index switch operation with a single linear target. + /// + /// Constructs an index switch operation for one qubit or qtensor value. + /// Each branch callback receives and returns a single value, avoiding + /// one-element ranges and vectors. + /// + /// @param arg Index argument. + /// @param target Initial argument for every index switch branch. + /// @param cases The individual switch cases. + /// @param caseBodies Functions that build the case bodies. + /// @param defaultBody Function that builds the default body. + /// @return The single result value. + /// + /// @par Example: + /// ```c++ + /// result = builder.qcoIndexSwitch( + /// arg, target, SmallVector{0}, + /// SmallVector>{ + /// [&](Value value) { return builder.x(value); }}, + /// [&](Value value) { return builder.z(value); }); + /// ``` Value qcoIndexSwitch(const std::variant& arg, Value target, ArrayRef cases, ArrayRef> caseBodies, function_ref defaultBody); - /** - * @brief Construct an scf.for operation - * - * @details - * Constructs an scf.for operation with the given loop boundaries and stepsize - * and a range of qubit and qtensor values for its iter args. Qubits that were - * extracted from a tensor that is used as an argument for this operation are - * automatically inserted before the operation is constructed. - * - * @param lowerbound Lower bound of the loop - * @param upperbound Upper bound of the loop - * @param step Step size of the loop - * @param initArgs Initial arguments for the iter args - * @param body Function that builds the body of the for operation - * @return ValueRange of the results - * - * @par Example: - * ```c++ - * builder.scfFor(lb, ub, step, initArgs, [&](Value iv, ValueRange iterArgs) - * -> SmallVector { - * auto [t0, q0] = builder.qtensorExtract(iterArgs[0], iv); - * auto q1 = builder.h(q0); - * auto insert = builder.qtensorInsert(q1, t0, iv); - * return {insert}; - * }); - * ``` - * ```mlir - * %t3 = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = %t0) - * -> (tensor<3x!qco.qubit>) { - * %t1, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> - * %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit - * %t2 = qtensor.insert %q1 into %t1[%iv] : tensor<3x!qco.qubit> - * scf.yield %t2 : tensor<3x!qco.qubit> - * } - * ``` - */ + /// Construct an scf.for operation + /// + /// Constructs an scf.for operation with the given loop boundaries and + /// stepsize and a range of qubit and qtensor values for its iter args. Qubits + /// that were extracted from a tensor that is used as an argument for this + /// operation are automatically inserted before the operation is constructed. + /// + /// @param lowerbound Lower bound of the loop + /// @param upperbound Upper bound of the loop + /// @param step Step size of the loop + /// @param initArgs Initial arguments for the iter args + /// @param body Function that builds the body of the for operation + /// @return ValueRange of the results + /// + /// @par Example: + /// ```c++ + /// builder.scfFor(lb, ub, step, initArgs, [&](Value iv, ValueRange iterArgs) + /// -> SmallVector { + /// auto [t0, q0] = builder.qtensorExtract(iterArgs[0], iv); + /// auto q1 = builder.h(q0); + /// auto insert = builder.qtensorInsert(q1, t0, iv); + /// return {insert}; + /// }); + /// ``` + /// ```mlir + /// %t3 = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = %t0) + /// -> (tensor<3x!qco.qubit>) { + /// %t1, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> + /// %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit + /// %t2 = qtensor.insert %q1 into %t1[%iv] : tensor<3x!qco.qubit> + /// scf.yield %t2 : tensor<3x!qco.qubit> + /// } + /// ``` ValueRange scfFor(const std::variant& lowerbound, const std::variant& upperbound, const std::variant& step, ValueRange initArgs, function_ref(Value, ValueRange)> body); - /** - * @brief Construct an scf.while operation - * - * @details - * Constructs an scf.while with a range of qubit and qtensor values for its - * iter args. Qubits that were extracted from a tensor that is used as an - * argument for this operation are automatically inserted before the operation - * is constructed. - * - * @param initArgs Arguments for the while loop - * @param beforeBody Function that builds the before body of the while - * operation - * @param afterBody Function that builds the after body of the while operation - * @return ValueRange of the results - * - * @par Example: - * ```c++ - * builder.scfWhile(initArgs, [&](ValueRange iterArgs) -> - * SmallVector { - * auto [q0, cond] = builder.measure(iterArgs[0]); - * builder.scfCondition(cond, q0); - * return {q0}; - * }, [&](ValueRange iterArgs) -> SmallVector { - * auto q0 = builder.h(iterArgs[0]); - * return {q0}; - * }); - * ``` - * ```mlir - * %q2 = scf.while (%arg0 = %q0): (!qco.qubit) -> (!qco.qubit) { - * %q1, %cond = qco.measure %arg0 : !qco.qubit - * scf.condition(%cond) %q1 : !qco.qubit - * } do { - * ^bb0(%arg0 : !qco.qubit): - * %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit - * scf.yield %q1 : !qco.qubit - * } - * ``` - */ + /// Construct an scf.while operation + /// + /// Constructs an scf.while with a range of qubit and qtensor values for its + /// iter args. Qubits that were extracted from a tensor that is used as an + /// argument for this operation are automatically inserted before the + /// operation is constructed. + /// + /// @param initArgs Arguments for the while loop + /// @param beforeBody Function that builds the before body of the while + /// operation + /// @param afterBody Function that builds the after body of the while + /// operation + /// @return ValueRange of the results + /// + /// @par Example: + /// ```c++ + /// builder.scfWhile(initArgs, [&](ValueRange iterArgs) -> + /// SmallVector { + /// auto [q0, cond] = builder.measure(iterArgs[0]); + /// builder.scfCondition(cond, q0); + /// return {q0}; + /// }, [&](ValueRange iterArgs) -> SmallVector { + /// auto q0 = builder.h(iterArgs[0]); + /// return {q0}; + /// }); + /// ``` + /// ```mlir + /// %q2 = scf.while (%arg0 = %q0): (!qco.qubit) -> (!qco.qubit) { + /// %q1, %cond = qco.measure %arg0 : !qco.qubit + /// scf.condition(%cond) %q1 : !qco.qubit + /// } do { + /// ^bb0(%arg0 : !qco.qubit): + /// %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit + /// scf.yield %q1 : !qco.qubit + /// } + /// ``` ValueRange scfWhile(ValueRange initArgs, function_ref(ValueRange)> beforeBody, function_ref(ValueRange)> afterBody); - /** - * @brief Construct an scf.condition operation with yielded values - * - * @param condition Condition for the condition operation - * @param yieldedValues ValueRange of the yieldedValues - * @return Reference to this builder for method chaining - * - * @par Example: - * ```c++ - * builder.scfCondition(condition, q0); - * ``` - * ```mlir - * scf.condition(%condition) %q0 : !qco.qubit - * ``` - */ + /// Construct an scf.condition operation with yielded values + /// + /// @param condition Condition for the condition operation + /// @param yieldedValues ValueRange of the yieldedValues + /// @return Reference to this builder for method chaining + /// + /// @par Example: + /// ```c++ + /// builder.scfCondition(condition, q0); + /// ``` + /// ```mlir + /// scf.condition(%condition) %q0 : !qco.qubit + /// ``` QCOProgramBuilder& scfCondition(Value condition, ValueRange yieldedValues); - /** - * @brief Construct an scf.condition operation conditioned on a classical bit - * - * @details Loads the classical bit from the given classical register at the - * given index and uses it as the condition of the condition operation. - * - * @param reg The CBit register - * @param index The index within the register to load the condition from - * @param yieldedValues ValueRange of the yielded values - * @return Reference to this builder for method chaining - */ + /// Construct an scf.condition operation conditioned on a classical bit + /// + /// Loads the classical bit from the given classical register at the + /// given index and uses it as the condition of the condition operation. + /// + /// @param reg The CBit register + /// @param index The index within the register to load the condition from + /// @param yieldedValues ValueRange of the yielded values + /// @return Reference to this builder for method chaining QCOProgramBuilder& scfCondition(Value reg, const std::variant& index, ValueRange yieldedValues); @@ -1886,60 +1787,51 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { // Finalization //===--------------------------------------------------------------------===// - /** - * @brief Finalize the program and return the constructed module - * - * @details - * Automatically deallocates all remaining valid qubits and tensors of qubits, - * adds a return statement with exit code 0 (indicating successful execution), - * and transfers ownership of the module to the caller. The builder should not - * be used after calling this method. - * - * @return OwningOpRef containing the constructed quantum program module - */ + /// Finalize the program and return the constructed module + /// + /// Automatically deallocates all remaining valid qubits and tensors of + /// qubits, adds a return statement with exit code 0 (indicating successful + /// execution), and transfers ownership of the module to the caller. The + /// builder should not be used after calling this method. + /// + /// @return OwningOpRef containing the constructed quantum program module OwningOpRef finalize(); - /** - * @brief Finalize the program with the given return values and return the - * constructed module - * @param returnValues Values representing the return values of the main - * function. - * - * @details - * Automatically deallocates all remaining valid qubits and tensors of qubits, - * adds a return statement with the given return values, and - * transfers ownership of the module to the caller. The builder should not - * be used after calling this method. - * - * The return values must have the types indicated by the function signature - * of the main function, which returns an `i64` by default and can be - * modified by passing different arguments to the `initialize()` method. - * - * @return OwningOpRef containing the constructed quantum program module - */ + /// Finalize the program with the given return values and return the + /// constructed module + /// @param returnValues Values representing the return values of the main + /// function. + /// + /// Automatically deallocates all remaining valid qubits and tensors of + /// qubits, adds a return statement with the given return values, and + /// transfers ownership of the module to the caller. The builder should not + /// be used after calling this method. + /// + /// The return values must have the types indicated by the function signature + /// of the main function, which returns an `i64` by default and can be + /// modified by passing different arguments to the `initialize()` method. + /// + /// @return OwningOpRef containing the constructed quantum program module OwningOpRef finalize(ValueRange returnValues); - /** - * @brief Convenience method for building quantum programs. - * @param context The MLIR context to use for building the program - * @param buildFunc A function that takes a reference to a QCOProgramBuilder - * and uses it to build the desired quantum program. The builder will be - * properly initialized before calling this function, and the resulting module - * will be finalized using the returned Values after this function completes. - * @return The module containing the quantum program built by buildFunc. - */ + /// Convenience method for building quantum programs. + /// @param context The MLIR context to use for building the program + /// @param buildFunc A function that takes a reference to a QCOProgramBuilder + /// and uses it to build the desired quantum program. The builder will be + /// properly initialized before calling this function, and the resulting + /// module will be finalized using the returned Values after this function + /// completes. + /// @return The module containing the quantum program built by buildFunc. static OwningOpRef build(MLIRContext* context, const function_ref(QCOProgramBuilder&)>& buildFunc); - /** - * @brief Convenience method for building quantum programs with one return - * value. - * @param context The MLIR context to use for building the program - * @param buildFunc A function that takes a reference to a QCOProgramBuilder - * and returns the single result value of the desired quantum program. - * @return The module containing the quantum program built by buildFunc. - */ + /// Convenience method for building quantum programs with one return + /// value. + /// @param context The MLIR context to use for building the program + /// @param buildFunc A function that takes a reference to a QCOProgramBuilder + /// and returns the single result value of the desired quantum program. + /// @return The module containing the quantum program built by buildFunc. static OwningOpRef build(MLIRContext* context, const function_ref& buildFunc); @@ -1948,7 +1840,7 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { enum class AllocationMode : uint8_t { Unset, Static, Dynamic }; MLIRContext* ctx{}; - Operation* module; + Operation* moduleOp_; /// Check if the builder has been finalized void checkFinalized() const; @@ -1957,18 +1849,14 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { // Linear Type Tracking Helpers //===--------------------------------------------------------------------===// - /** - * @brief Validate that a qubit value is valid and unconsumed - * @param qubit Qubit value to validate - * @throws Aborts if qubit is not tracked (consumed or never created) - */ + /// Validate that a qubit value is valid and unconsumed + /// @param qubit Qubit value to validate + /// @throws Aborts if qubit is not tracked (consumed or never created) void validateQubitValue(Value qubit) const; - /** - * @brief Update tracking when an operation consumes and produces a qubit - * @param inputQubit Input qubit being consumed (must be valid) - * @param outputQubit New output qubit being produced - */ + /// Update tracking when an operation consumes and produces a qubit + /// @param inputQubit Input qubit being consumed (must be valid) + /// @param outputQubit New output qubit being produced void updateQubitTracking(Value inputQubit, Value outputQubit); /// Count unique tensors @@ -1989,65 +1877,54 @@ class QCOProgramBuilder final : public ImplicitLocOpBuilder { /// is removed and the new output is added. DenseSet validQubits; - /** - * @brief Validate that a tensor value is valid and unconsumed. This also - * checks if the tensor is one-dimensional and contains !qco.qubit as its - * values - * @param tensor Tensor value to validate - * @throws Aborts if tensor is not tracked (consumed or never created) - */ + /// Validate that a tensor value is valid and unconsumed. This also + /// checks if the tensor is one-dimensional and contains !qco.qubit as its + /// values + /// @param tensor Tensor value to validate + /// @throws Aborts if tensor is not tracked (consumed or never created) void validateTensorValue(Value tensor) const; - /** - * @brief Update tracking when an operation consumes and produces a tensor - * @param inputTensor Input tensor being consumed (must be valid) - * @param outputTensor New output tensor being produced - */ + /// Update tracking when an operation consumes and produces a tensor + /// @param inputTensor Input tensor being consumed (must be valid) + /// @param outputTensor New output tensor being produced void updateTensorTracking(Value inputTensor, Value outputTensor); - /** - * @brief Prepares initial arguments for operations by re-inserting extracted - * qubits into their tensors - * - * @details For each tensor in @p initArgs, any qubits extracted from it that - * are not also present in @p initArgs are inserted back. The latest tensor - * values after inserting the qubits are returned. Qubit values are returned - * without modifications. - * - * @param initArgs ValueRange of the initial values - * @return SmallVector of the updated values of the initial values. - */ + /// Dispose of every live linear value in the current function. + void disposeLinearValues(); + + /// Prepares initial arguments for operations by re-inserting extracted + /// qubits into their tensors + /// + /// For each tensor in @p initArgs, any qubits extracted from it that + /// are not also present in @p initArgs are inserted back. The latest tensor + /// values after inserting the qubits are returned. Qubit values are returned + /// without modifications. + /// + /// @param initArgs ValueRange of the initial values + /// @return SmallVector of the updated values of the initial values. SmallVector prepareInitArgs(ValueRange initArgs); - /** - * @brief Prepare one initial argument by re-inserting extracted qubits into - * its tensor, if necessary. - * @param initArg Initial value - * @return Updated initial value - */ + /// Prepare one initial argument by re-inserting extracted qubits into + /// its tensor, if necessary. + /// @param initArg Initial value + /// @return Updated initial value Value prepareInitArg(Value initArg); Value prepareInitArg(Value initArg, const DenseSet* initQubits); - /** - * @brief Update linear-value tracking for one replaced value - * @param oldValue The old value to be replaced - * @param newValue The new value to be tracked - */ + /// Update linear-value tracking for one replaced value + /// @param oldValue The old value to be replaced + /// @param newValue The new value to be tracked void updateQubitValueTracking(Value oldValue, Value newValue); - /** - * @brief Update the qubit tracking of the old values with the new values - * @param oldValues The old values to be replaced - * @param newValues The new values to be tracked - */ + /// Update the qubit tracking of the old values with the new values + /// @param oldValues The old values to be replaced + /// @param newValues The new values to be tracked void updateQubitValueTracking(ValueRange oldValues, ValueRange newValues); - /** - * @brief Check if every value is either a qubit or a tensor of qubits - * @param values The values that are checked - * @throws Abort if a value is neither a qubit nor a tensor of qubits - */ + /// Check if every value is either a qubit or a tensor of qubits + /// @param values The values that are checked + /// @throws Abort if a value is neither a qubit nor a tensor of qubits static void checkQubitType(ValueRange values); struct TensorDenseMapInfo { diff --git a/mlir/include/mlir/Dialect/QCO/IR/QCOOps.h b/mlir/include/mlir/Dialect/QCO/IR/QCOOps.h index 99d5bb6fc0..d698be870c 100644 --- a/mlir/include/mlir/Dialect/QCO/IR/QCOOps.h +++ b/mlir/include/mlir/Dialect/QCO/IR/QCOOps.h @@ -25,9 +25,12 @@ #include "mlir/Dialect/QCO/IR/QCOInterfaces.h" #include +#include +#include #include #include +#include #include #define GET_OP_CLASSES diff --git a/mlir/include/mlir/Dialect/QCO/IR/QCOOps.td b/mlir/include/mlir/Dialect/QCO/IR/QCOOps.td index f9a8971dfd..9d2842ecd7 100644 --- a/mlir/include/mlir/Dialect/QCO/IR/QCOOps.td +++ b/mlir/include/mlir/Dialect/QCO/IR/QCOOps.td @@ -15,6 +15,8 @@ include "mlir/Dialect/QCO/IR/QCOTypes.td" include "mlir/IR/EnumAttr.td" include "mlir/IR/OpBase.td" +include "mlir/IR/SymbolInterfaces.td" +include "mlir/Interfaces/CallInterfaces.td" include "mlir/Interfaces/ControlFlowInterfaces.td" include "mlir/Interfaces/InferTypeOpInterface.td" include "mlir/Interfaces/SideEffectInterfaces.td" @@ -1133,6 +1135,78 @@ def BarrierOp : QCOOp<"barrier", traits = [UnitaryOpInterface, Pure]> { let hasCanonicalizer = 1; } +def CallOp + : QCOOp<"call", + traits = [CallOpInterface, UnitaryOpInterface, + DeclareOpInterfaceMethods, Pure]> { + let summary = "Call a unitary QCO function"; + let description = [{ + Calls a private `func.func` marked with `mqt.unitary`. Parameters precede + qubit operands. Each qubit result continues the corresponding qubit input. + + Example: + ```mlir + %q1 = qco.call @rotate(%theta, %q0) + : (f64, !qco.qubit) -> !qco.qubit + ``` + }]; + + let arguments = (ins FlatSymbolRefAttr:$callee, Variadic:$operands, + OptionalAttr:$arg_attrs, + OptionalAttr:$res_attrs); + let results = (outs Variadic:$qubits_out); + let assemblyFormat = [{ + $callee `(` $operands `)` attr-dict `:` + functional-type($operands, $qubits_out) + }]; + + let builders = [OpBuilder<(ins "FlatSymbolRefAttr":$callee, + "ValueRange":$operands)>]; + + let hasVerifier = 1; + + let extraClassDeclaration = [{ + size_t getNumQubits() { return getQubitsOut().size(); } + size_t getNumTargets() { return getNumQubits(); } + static size_t getNumControls() { return 0; } + Value getInputQubit(size_t i) { return getInputTarget(i); } + OperandRange getInputQubits(); + Value getOutputQubit(size_t i) { return getOutputTarget(i); } + ResultRange getOutputQubits() { return getQubitsOut(); } + Value getInputTarget(size_t i) { return getInputQubits()[i]; } + OperandRange getInputTargets() { return getInputQubits(); } + Value getOutputTarget(size_t i) { return getQubitsOut()[i]; } + ResultRange getOutputTargets() { return getQubitsOut(); } + static Value getInputControl(size_t i) { + llvm::reportFatalUsageError("CallOp does not have controls"); + } + static OperandRange getInputControls() { return {nullptr, 0}; } + static Value getOutputControl(size_t i) { + llvm::reportFatalUsageError("CallOp does not have controls"); + } + static ResultRange getOutputControls() { return {nullptr, 0}; } + Value getInputForOutput(Value output); + Value getOutputForInput(Value input); + size_t getNumParams(); + Value getParameter(size_t i) { return getParameters()[i]; } + OperandRange getParameters(); + StringRef getBaseSymbol() { return getCallee(); } + static bool hasCompileTimeKnownUnitaryMatrix() { return false; } + static std::optional getUnitaryMatrix() { + return std::nullopt; + } + + ::mlir::Operation::operand_range getArgOperands() { return getOperands(); } + MutableOperandRange getArgOperandsMutable() { return getOperandsMutable(); } + ::mlir::CallInterfaceCallable getCallableForCallee() { + return getCalleeAttr(); + } + void setCalleeFromCallable(::mlir::CallInterfaceCallable callee) { + setCalleeAttr(cast(cast(callee))); + } + }]; +} + //===----------------------------------------------------------------------===// // Modifiers //===----------------------------------------------------------------------===// diff --git a/mlir/include/mlir/Dialect/QCO/Utils/FunctionUtils.h b/mlir/include/mlir/Dialect/QCO/Utils/FunctionUtils.h new file mode 100644 index 0000000000..e0ef07c132 --- /dev/null +++ b/mlir/include/mlir/Dialect/QCO/Utils/FunctionUtils.h @@ -0,0 +1,24 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#pragma once + +#include +#include +#include + +namespace mlir::qco { +/// Return the qubit argument continued by @p value. +/// +/// QCO functions return one trailing qubit for every qubit argument, in +/// qubit-argument order. Generic calls are followed only through that ABI. +[[nodiscard]] FailureOr traceQubitArgument(func::FuncOp function, + Value value); +} // namespace mlir::qco diff --git a/mlir/include/mlir/Dialect/QCO/Utils/WireIterator.h b/mlir/include/mlir/Dialect/QCO/Utils/WireIterator.h index 8813bc774a..640f597ca4 100644 --- a/mlir/include/mlir/Dialect/QCO/Utils/WireIterator.h +++ b/mlir/include/mlir/Dialect/QCO/Utils/WireIterator.h @@ -10,59 +10,20 @@ #pragma once -#include -#include -#include #include #include -#include #include -#include #include namespace mlir::qco { -/// Resolves how qubits flow across call boundaries. -/// -/// The mapping follows each qubit argument through the callee instead of -/// assuming positional correspondence. Results are cached per callee. Mapping -/// fails for declarations, recursion, and non-straight-line bodies. -class CallQubitMapping { -public: - /// Gets the result continuing @p operand's wire. - /// - /// Returns a null value when the callee keeps the qubit and failure when the - /// correspondence cannot be derived. - [[nodiscard]] FailureOr getResultForOperand(func::CallOp callOp, - Value operand); - - /// Clears all cached correspondence after a callee is changed or erased. - void invalidate(); - -private: - friend class WireIterator; - - // Marks a qubit argument that never reaches a result. - static constexpr int64_t KEPT = -1; - - // Returns each qubit argument's call-result index, or KEPT. - FailureOr> mappingFor(func::CallOp callOp); - - // Derives a mapping by threading every qubit argument through the callee. - FailureOr> computeMapping(func::FuncOp callee); - - // Gets the call operand feeding a result's wire. - FailureOr getOperandForResult(func::CallOp callOp, Value result); - - DenseMap> cache; - DenseSet inProgress; -}; - /// A bidirectional iterator over the def-use chain of a qubit wire. /// /// The iterator follows the flow of a qubit through a sequence of quantum /// operations while respecting the semantics of each operation. +/// Unitary calls preserve the wire without entering the callee. Generic +/// `func.call` operations end traversal; clients manage callee traversal. class [[nodiscard]] WireIterator { public: using iterator_category = std::bidirectional_iterator_tag; @@ -70,13 +31,11 @@ class [[nodiscard]] WireIterator { using value_type = Operation*; /// Construct a dead-end sentinel wire-iterator. - WireIterator() - : mapping_(nullptr), op_(nullptr), qubit_(nullptr), - pos_(Position::PastTail) {} + WireIterator() : op_(nullptr), qubit_(nullptr), pos_(Position::PastTail) {} /// Construct a wire iterator pointing at the defining op of a qubit value. - explicit WireIterator(Value qubit, CallQubitMapping* mapping = nullptr) - : mapping_(mapping), op_(qubit.getDefiningOp()), qubit_(qubit) { + explicit WireIterator(Value qubit) + : op_(qubit.getDefiningOp()), qubit_(qubit) { if (op_ == nullptr || isHead(op_)) { pos_ = Position::Head; } else if (isTail(op_)) { @@ -127,8 +86,6 @@ class [[nodiscard]] WireIterator { } private: - friend class CallQubitMapping; - /// Labels the position on the wire. enum class Position : uint8_t { BeforeHead, Head, Between, Tail, PastTail }; @@ -144,18 +101,9 @@ class [[nodiscard]] WireIterator { // Moves to the previous operation on the qubit wire. void backward(); - // Resolves the call result continuing an operand's wire. - FailureOr resultForOperand(func::CallOp callOp, Value operand) const; - - // Resolves the call operand feeding a result's wire. - [[nodiscard]] Value operandForResult(func::CallOp callOp, Value result) const; - - // Null means that each call query uses a fresh mapping. - CallQubitMapping* mapping_; Operation* op_; Value qubit_; Position pos_; - bool mappingFailed_ = false; }; /// Categorizes the current traversal direction. diff --git a/mlir/include/mlir/Dialect/QTensor/Utils/TensorIterator.h b/mlir/include/mlir/Dialect/QTensor/Utils/TensorIterator.h index 7f4b97ac80..9a063cd63a 100644 --- a/mlir/include/mlir/Dialect/QTensor/Utils/TensorIterator.h +++ b/mlir/include/mlir/Dialect/QTensor/Utils/TensorIterator.h @@ -10,24 +10,19 @@ #pragma once -#include -#include -#include -#include -#include #include #include #include #include #include -#include -#include #include namespace mlir::qtensor { /// A bidirectional iterator traversing the tensor chain. +/// +/// `func.call` operations end traversal; clients manage callee traversal. class [[nodiscard]] TensorIterator { public: using iterator_category = std::bidirectional_iterator_tag; @@ -93,34 +88,4 @@ class [[nodiscard]] TensorIterator { bool isSentinel_; }; -/// Resolves how qubit tensors flow across call boundaries. -/// -/// The mapping follows each tensor argument through the callee instead of -/// assuming positional correspondence. Results are cached per callee. Mapping -/// fails for declarations, recursion, and non-straight-line bodies. -class CallTensorMapping { -public: - /// Gets the result continuing @p operand's tensor chain. - /// - /// Returns a null value when the callee keeps the tensor and failure when the - /// correspondence cannot be derived. - [[nodiscard]] FailureOr getResultForOperand(func::CallOp callOp, - Value operand); - -private: - // Marks a tensor argument that never reaches a result. - static constexpr int64_t KEPT = -1; - - // Returns each tensor argument's call-result index, or KEPT. - FailureOr> mappingFor(func::CallOp callOp); - - // Derives a mapping by threading every tensor argument through the callee. - FailureOr> computeMapping(func::FuncOp callee); - - // Follows an argument to a return operand, hopping over calls. - FailureOr threadToResult(Value arg, func::ReturnOp returnOp); - - DenseMap> cache; - DenseSet inProgress; -}; } // namespace mlir::qtensor diff --git a/mlir/lib/Conversion/QCOToQC/CMakeLists.txt b/mlir/lib/Conversion/QCOToQC/CMakeLists.txt index 2fad50932a..fcf03b6330 100644 --- a/mlir/lib/Conversion/QCOToQC/CMakeLists.txt +++ b/mlir/lib/Conversion/QCOToQC/CMakeLists.txt @@ -17,6 +17,7 @@ add_mlir_conversion_library( MLIRCBitDialect MLIRQCDialect MLIRQCODialect + MLIRQCOUtils MLIRQTensorDialect MLIRArithDialect MLIRFuncDialect diff --git a/mlir/lib/Conversion/QCOToQC/QCOToQC.cpp b/mlir/lib/Conversion/QCOToQC/QCOToQC.cpp index b6b519b269..a9e4a78636 100644 --- a/mlir/lib/Conversion/QCOToQC/QCOToQC.cpp +++ b/mlir/lib/Conversion/QCOToQC/QCOToQC.cpp @@ -12,14 +12,17 @@ #include "mlir/Conversion/ConversionUtils.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" +#include "mlir/Dialect/MQT/IR/MQTDialect.h" #include "mlir/Dialect/QC/IR/QCDialect.h" #include "mlir/Dialect/QC/IR/QCOps.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/IR/QCOOps.h" +#include "mlir/Dialect/QCO/Utils/FunctionUtils.h" #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include +#include #include #include #include @@ -30,6 +33,7 @@ #include #include #include +#include #include #include #include @@ -47,25 +51,25 @@ using namespace qc; namespace { -/** @brief Qubit allocation mode */ +/// Qubit allocation mode enum class AllocationMode : std::uint8_t { Unset, //!< No allocation mode has been established yet. Static, //!< The module uses static qubit allocation. Dynamic //!< The module uses dynamic qubit allocation. }; -/** - * @brief State object for tracking qubit allocation mode. - * - * @details - * Used to track whether a function uses static or dynamic qubit allocation. - * This is used to determine whether to convert `qco.sink` to `qc.dealloc` (for - * dynamic qubits) or simply erase it (for static qubits). This is also used to - * catch cases of mixed allocation modes being used, which is not supported. - */ +/// State object for tracking qubit allocation mode. +/// +/// Used to track whether a function uses static or dynamic qubit allocation. +/// This is used to determine whether to convert `qco.sink` to `qc.dealloc` (for +/// dynamic qubits) or simply erase it (for static qubits). This is also used to +/// catch cases of mixed allocation modes being used, which is not supported. struct LoweringState { /// Per-region map from a register's indices to its loaded qubit values. DenseMap>> qubitValues; + /// Original qubit argument positions, retained while signatures are + /// rewritten. + DenseMap> qubitArguments; /// The qubit allocation mode used in the module AllocationMode allocationMode = AllocationMode::Unset; @@ -84,14 +88,11 @@ struct LoweringState { } }; -/** - * @brief Base class for conversion patterns that need access to lowering state - * - * @details - * Extends OpConversionPattern to provide access to a shared LoweringState - * object, which is used to track the allocation mode of the module. - * @tparam OpType The QCO operation type to be converted. - */ +/// Base class for conversion patterns that need access to lowering state +/// +/// Extends OpConversionPattern to provide access to a shared LoweringState +/// object, which is used to track the allocation mode of the module. +/// @tparam OpType The QCO operation type to be converted. template class StatefulOpConversionPattern : public OpConversionPattern { @@ -108,19 +109,17 @@ class StatefulOpConversionPattern : public OpConversionPattern { }; } // namespace -/** - * @brief Moves the operations from one region into another. - * - * @details Moves the operations from the source region into the target region. - * The target region replaces the uses of the old block arguments with the - * @p replacementValues and erases the unused block arguments. - * - * @param sourceRegion Source region where the operations are moved from - * @param targetRegion Target region where the operations are moved to - * @param offset Offset to the arguments that are dropped - * @param replacementValues Values to replace the uses of the arguments - * @param rewriter PatternRewriter of the current conversion pass - */ +/// Moves the operations from one region into another. +/// +/// Moves the operations from the source region into the target region. +/// The target region replaces the uses of the old block arguments with the +/// @p replacementValues and erases the unused block arguments. +/// +/// @param sourceRegion Source region where the operations are moved from +/// @param targetRegion Target region where the operations are moved to +/// @param offset Offset to the arguments that are dropped +/// @param replacementValues Values to replace the uses of the arguments +/// @param rewriter PatternRewriter of the current conversion pass static void inlineRegion(Region& sourceRegion, Region& targetRegion, unsigned int offset, ValueRange replacementValues, ConversionPatternRewriter& rewriter) { @@ -214,21 +213,18 @@ combineConvertedResults(TypeRange originalTypes, ValueRange classicalResults, namespace { -/** - * @brief Type converter for QCO-to-QC conversion - * - * @details - * Handles type conversion between the QCO and QC dialects. - * The primary conversion is from !qco.qubit to !qc.qubit, which - * represents the semantic shift from value types to reference types. - * - * Qubit tensor types preserve their shape during conversion: a statically - * shaped `tensor` becomes `memref`, while a - * dynamically shaped `tensor` becomes `memref`. - * - * Other types (integers, booleans, etc.) pass through unchanged via - * the identity conversion. - */ +/// Type converter for QCO-to-QC conversion +/// +/// Handles type conversion between the QCO and QC dialects. +/// The primary conversion is from !qco.qubit to !qc.qubit, which +/// represents the semantic shift from value types to reference types. +/// +/// Qubit tensor types preserve their shape during conversion: a statically +/// shaped `tensor` becomes `memref`, while a +/// dynamically shaped `tensor` becomes `memref`. +/// +/// Other types (integers, booleans, etc.) pass through unchanged via +/// the identity conversion. class QCOToQCTypeConverter final : public TypeConverter { public: explicit QCOToQCTypeConverter(MLIRContext* ctx) { @@ -249,18 +245,199 @@ class QCOToQCTypeConverter final : public TypeConverter { } }; -/** - * @brief Converts qtensor.alloc to memref.alloc - * - * @par Example: - * ```mlir - * %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - * ``` - * is converted to - * ```mlir - * %memref = memref.alloc(%c3) : memref<3x!qc.qubit> - * ``` - */ +} // namespace + +[[nodiscard]] static LogicalResult +collectFunctionQubitArguments(ModuleOp moduleOp, LoweringState& state) { + for (auto function : moduleOp.getOps()) { + auto& qubitArguments = state.qubitArguments[function]; + for (auto [index, type] : llvm::enumerate(function.getArgumentTypes())) { + if (isa(type)) { + qubitArguments.emplace_back(index); + } + } + if (qubitArguments.empty()) { + continue; + } + if (function.getNumResults() < qubitArguments.size() || + llvm::any_of(function.getResultTypes().take_back(qubitArguments.size()), + [](Type type) { return !isa(type); })) { + return function.emitOpError() + << "must return one trailing qubit for each qubit argument"; + } + const auto firstQubitResult = + function.getNumResults() - qubitArguments.size(); + for (unsigned index = firstQubitResult; index < function.getNumResults(); + ++index) { + if (auto attrs = function.getResultAttrDict(index); + attrs && !attrs.empty()) { + return function.emitOpError( + "cannot preserve attributes on pass-through qubit results in QC"); + } + } + if (function.isDeclaration()) { + continue; + } + if (!function.getBody().hasOneBlock()) { + return function.emitOpError() + << "with qubit arguments must have one outer block"; + } + auto returnOp = dyn_cast(function.getBody().front().back()); + if (!returnOp) { + return function.emitOpError("must terminate with func.return"); + } + auto returnedQubits = + returnOp.getOperands().take_back(qubitArguments.size()); + for (auto [argument, value] : + llvm::zip_equal(qubitArguments, returnedQubits)) { + auto origin = qco::traceQubitArgument(function, value); + if (failed(origin) || *origin != argument) { + return function.emitOpError() + << "must return its qubit arguments positionally"; + } + } + } + return success(); +} + +namespace { + +struct ConvertFuncOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(func::FuncOp op, OpAdaptor, + ConversionPatternRewriter& rewriter) const override { + TypeConverter::SignatureConversion signature(op.getNumArguments()); + if (failed(getTypeConverter()->convertSignatureArgs(op.getArgumentTypes(), + signature))) { + return failure(); + } + SmallVector inputs; + if (failed( + getTypeConverter()->convertTypes(op.getArgumentTypes(), inputs))) { + return failure(); + } + + const auto& qubitArguments = getState().qubitArguments[op]; + const auto firstQubitResult = op.getNumResults() - qubitArguments.size(); + SmallVector results; + if (failed(getTypeConverter()->convertTypes( + op.getResultTypes().take_front(firstQubitResult), results))) { + return failure(); + } + SmallVector resultAttrs; + for (unsigned index = 0; index < firstQubitResult; ++index) { + resultAttrs.emplace_back(op.getResultAttrDict(index)); + } + + rewriter.modifyOpInPlace(op, [&] { + op.setType(rewriter.getFunctionType(inputs, results)); + function_interface_impl::setAllResultAttrDicts(op, resultAttrs); + }); + if (!op.isExternal() && + failed(rewriter.convertRegionTypes(&op.getBody(), *getTypeConverter(), + &signature))) { + return failure(); + } + return success(); + } +}; + +struct ConvertFuncReturnOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(func::ReturnOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + auto function = op->getParentOfType(); + const auto numQubitArguments = getState().qubitArguments[function].size(); + rewriter.replaceOpWithNewOp( + op, adaptor.getOperands().drop_back(numQubitArguments)); + return success(); + } +}; + +struct ConvertFuncCallOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(func::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + auto callee = SymbolTable::lookupNearestSymbolFrom( + op, op.getCalleeAttr()); + if (!callee) { + return rewriter.notifyMatchFailure(op, "callee is not defined"); + } + const auto& qubitArguments = getState().qubitArguments[callee]; + const auto firstQubitResult = op.getNumResults() - qubitArguments.size(); + auto resultAttrs = op.getResAttrsAttr(); + if (resultAttrs && + llvm::any_of(resultAttrs.getValue().take_back(qubitArguments.size()), + [](Attribute attr) { + return !cast(attr).empty(); + })) { + return op.emitOpError( + "cannot preserve attributes on pass-through qubit results in QC"); + } + + SmallVector keptResultTypes(op.getResultTypes()); + keptResultTypes.resize(firstQubitResult); + SmallVector resultTypes; + if (failed( + getTypeConverter()->convertTypes(keptResultTypes, resultTypes))) { + return failure(); + } + auto call = func::CallOp::create(rewriter, op.getLoc(), op.getCallee(), + resultTypes, adaptor.getOperands()); + call->setAttrs(op->getAttrs()); + if (resultAttrs) { + call.setResAttrsAttr(rewriter.getArrayAttr( + resultAttrs.getValue().take_front(firstQubitResult))); + } + + SmallVector replacements; + llvm::append_range(replacements, call.getResults()); + for (const auto argument : qubitArguments) { + replacements.emplace_back(adaptor.getOperands()[argument]); + } + rewriter.replaceOp(op, replacements); + return success(); + } +}; + +struct ConvertQCOCallOp final : OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(qco::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + if (auto attrs = op.getResAttrsAttr(); + attrs && llvm::any_of(attrs, [](Attribute attr) { + return !cast(attr).empty(); + })) { + return op.emitOpError( + "cannot preserve unitary call result attributes in QC"); + } + auto call = qc::CallOp::create(rewriter, op.getLoc(), op.getCalleeAttr(), + adaptor.getOperands()); + call->setAttrs(op->getAttrs()); + call.removeResAttrsAttr(); + rewriter.replaceOp(op, adaptor.getOperands().take_back(op.getNumResults())); + return success(); + } +}; + +/// Converts qtensor.alloc to memref.alloc +/// +/// @par Example: +/// ```mlir +/// %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> +/// ``` +/// is converted to +/// ```mlir +/// %memref = memref.alloc(%c3) : memref<3x!qc.qubit> +/// ``` struct ConvertQTensorAllocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -291,18 +468,16 @@ struct ConvertQTensorAllocOp final } }; -/** - * @brief Converts qtensor.extract to memref.load - * - * @par Example: - * ```mlir - * %tensor_out, %q = qtensor.extract %tensor_in[%c0]: tensor<3x!qco.qubit> - * ``` - * is converted to - * ```mlir - * %q = memref.load %memref[%c0] : memref<3x!qc.qubit> - * ``` - */ +/// Converts qtensor.extract to memref.load +/// +/// @par Example: +/// ```mlir +/// %tensor_out, %q = qtensor.extract %tensor_in[%c0]: tensor<3x!qco.qubit> +/// ``` +/// is converted to +/// ```mlir +/// %q = memref.load %memref[%c0] : memref<3x!qc.qubit> +/// ``` struct ConvertQTensorExtractOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -326,7 +501,7 @@ struct ConvertQTensorExtractOp final } }; -/** Converts qtensor.insert to an in-place memref.store. */ +/// Converts qtensor.insert to an in-place memref.store. struct ConvertQTensorInsertOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -354,18 +529,16 @@ struct ConvertQTensorInsertOp final } }; -/** - * @brief Converts qtensor.dealloc to memref.dealloc - * - * @par Example: - * ```mlir - * qtensor.dealloc %tensor : tensor<3x!qco.qubit> - * ``` - * is converted to - * ```mlir - * memref.dealloc %memref : memref<3x!qc.qubit> - * ``` - */ +/// Converts qtensor.dealloc to memref.dealloc +/// +/// @par Example: +/// ```mlir +/// qtensor.dealloc %tensor : tensor<3x!qco.qubit> +/// ``` +/// is converted to +/// ```mlir +/// memref.dealloc %memref : memref<3x!qc.qubit> +/// ``` struct ConvertQTensorDeallocOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -382,32 +555,29 @@ template { using OpConversionPattern::OpConversionPattern; - /** - * @brief Generic QCO gate conversion helper (value semantics -> reference). - * - * @details - * This helper relies on a strict operand ordering contract provided by the - * dialect conversion framework: - * - `adaptor.getOperands()` is expected to be ordered as - * `targets...` followed by `parameters...`. - * - The first @p NumTargets operands are the (type-converted) QC target - * qubits. - * - The remaining @p NumParams operands are the gate parameters. - * - * `matchAndRewrite` passes the full adapted operand list to `createGate`, - * which forwards the first @p NumTargets values (converted targets) and the - * following @p NumParams values (parameters, unchanged type through the - * converter) to `QCOpType::create(...)`. It then replaces the original QCO op - * with the created QC targets via `rewriter.replaceOp(op, qcTargets)`. - * - * The values of @p NumTargets and @p NumParams are compile-time constants and - * define this contract for each instantiation. - * - * @see ConvertQCOGateToQC - * @see createGate - * @see matchAndRewrite - * @see addGatePattern - */ + /// Generic QCO gate conversion helper (value semantics -> reference). + /// + /// This helper relies on a strict operand ordering contract provided by the + /// dialect conversion framework: + /// - `adaptor.getOperands()` is expected to be ordered as + /// `targets...` followed by `parameters...`. + /// - The first @p NumTargets operands are the (type-converted) QC target + /// qubits. + /// - The remaining @p NumParams operands are the gate parameters. + /// + /// `matchAndRewrite` passes the full adapted operand list to `createGate`, + /// which forwards the first @p NumTargets values (converted targets) and the + /// following @p NumParams values (parameters, unchanged type through the + /// converter) to `QCOpType::create(...)`. It then replaces the original QCO + /// op with the created QC targets via `rewriter.replaceOp(op, qcTargets)`. + /// + /// The values of @p NumTargets and @p NumParams are compile-time constants + /// and define this contract for each instantiation. + /// + /// @see ConvertQCOGateToQC + /// @see createGate + /// @see matchAndRewrite + /// @see addGatePattern template static void createGate(ConversionPatternRewriter& rewriter, Location loc, ValueRange qcOperands, @@ -433,7 +603,7 @@ struct ConvertQCOGateToQC final : OpConversionPattern { } }; -/** Converts a variadic dense qco.unitary to its reference-semantics form. */ +/// Converts a variadic dense qco.unitary to its reference-semantics form. struct ConvertQCOUnitaryOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -459,18 +629,16 @@ static void addGatePattern(RewritePatternSet& patterns, namespace { -/** - * @brief Converts qco.alloc to qc.alloc - * - * @par Example: - * ```mlir - * %q = qco.alloc : !qco.qubit - * ``` - * is converted to - * ```mlir - * %q = qc.alloc : !qc.qubit - * ``` - */ +/// Converts qco.alloc to qc.alloc +/// +/// @par Example: +/// ```mlir +/// %q = qco.alloc : !qco.qubit +/// ``` +/// is converted to +/// ```mlir +/// %q = qc.alloc : !qc.qubit +/// ``` struct ConvertQCOAllocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -489,26 +657,23 @@ struct ConvertQCOAllocOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qco.sink to qc.dealloc. - * - * @details - * In QCO, qubits have value/linear semantics and must be consumed explicitly - * (via `qco.sink`). In QC, qubits have reference semantics; for dynamic qubits - * we materialize this end-of-lifetime as `qc.dealloc`. Static qubits do not - * need explicit deallocation, so we simply erase the `qco.sink` operation. - * - * The OpAdaptor automatically provides the type-converted qubit operand - * (`!qc.qubit` instead of `!qco.qubit`), so we simply pass it through to the - * new operation when needed. - * - * Example transformation: - * ```mlir - * qco.sink %q_qco : !qco.qubit - * // becomes: - * qc.dealloc %q_qc : !qc.qubit - * ``` - */ +/// Converts qco.sink to qc.dealloc. +/// +/// In QCO, qubits have value/linear semantics and must be consumed explicitly +/// (via `qco.sink`). In QC, qubits have reference semantics; for dynamic qubits +/// we materialize this end-of-lifetime as `qc.dealloc`. Static qubits do not +/// need explicit deallocation, so we simply erase the `qco.sink` operation. +/// +/// The OpAdaptor automatically provides the type-converted qubit operand +/// (`!qc.qubit` instead of `!qco.qubit`), so we simply pass it through to the +/// new operation when needed. +/// +/// Example transformation: +/// ```mlir +/// qco.sink %q_qco : !qco.qubit +/// // becomes: +/// qc.dealloc %q_qc : !qc.qubit +/// ``` struct ConvertQCOSinkOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -530,21 +695,18 @@ struct ConvertQCOSinkOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qco.static to qc.static - * - * @details - * Static qubits represent references to hardware-mapped or fixed-position - * qubits identified by an index. The conversion preserves the index attribute - * and creates the corresponding qc.static operation. - * - * Example transformation: - * ```mlir - * %q0 = qco.static 0 : !qco.qubit - * // becomes: - * %q = qc.static 0 : !qc.qubit - * ``` - */ +/// Converts qco.static to qc.static +/// +/// Static qubits represent references to hardware-mapped or fixed-position +/// qubits identified by an index. The conversion preserves the index attribute +/// and creates the corresponding qc.static operation. +/// +/// Example transformation: +/// ```mlir +/// %q0 = qco.static 0 : !qco.qubit +/// // becomes: +/// %q = qc.static 0 : !qc.qubit +/// ``` struct ConvertQCOStaticOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -562,30 +724,28 @@ struct ConvertQCOStaticOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qco.measure to qc.measure - * - * @details - * Measurement demonstrates the key semantic difference between the dialects: - * - QCO (value semantics): Consumes input qubit, returns both output qubit - * and classical bit result - * - QC (reference semantics): Measures qubit in-place, returns only the - * classical bit result - * - * The OpAdaptor provides the input qubit already converted to !qc.qubit. - * Since QC operations are in-place, we return the same qubit reference - * alongside the measurement bit. MLIR's conversion infrastructure automatically - * routes subsequent uses of the QCO output qubit to this QC reference. - * - * @par Example: - * ```mlir - * %q_out, %c = qco.measure %q_in : !qco.qubit - * ``` - * is converted to - * ```mlir - * %c = qc.measure %q : !qc.qubit -> i1 - * ``` - */ +/// Converts qco.measure to qc.measure +/// +/// Measurement demonstrates the key semantic difference between the dialects: +/// - QCO (value semantics): Consumes input qubit, returns both output qubit +/// and classical bit result +/// - QC (reference semantics): Measures qubit in-place, returns only the +/// classical bit result +/// +/// The OpAdaptor provides the input qubit already converted to !qc.qubit. +/// Since QC operations are in-place, we return the same qubit reference +/// alongside the measurement bit. MLIR's conversion infrastructure +/// automatically routes subsequent uses of the QCO output qubit to this QC +/// reference. +/// +/// @par Example: +/// ```mlir +/// %q_out, %c = qco.measure %q_in : !qco.qubit +/// ``` +/// is converted to +/// ```mlir +/// %c = qc.measure %q : !qc.qubit -> i1 +/// ``` struct ConvertQCOMeasureOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -607,27 +767,24 @@ struct ConvertQCOMeasureOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.reset to qc.reset - * - * @details - * Reset operations force a qubit to the |0⟩ state: - * - QCO (value semantics): Consumes input qubit, returns reset output qubit - * - QC (reference semantics): Resets qubit in-place, no result value - * - * The OpAdaptor provides the input qubit already converted to !qc.qubit. - * Since QC's reset is in-place, we return the same qubit reference. - * MLIR's conversion infrastructure automatically routes subsequent uses of - * the QCO output qubit to this QC reference. - * - * Example transformation: - * ```mlir - * %q_out = qco.reset %q_in : !qco.qubit -> !qco.qubit - * // becomes: - * qc.reset %q : !qc.qubit - * // %q_out uses are replaced with %q (the adaptor-converted input) - * ``` - */ +/// Converts qco.reset to qc.reset +/// +/// Reset operations force a qubit to the |0⟩ state: +/// - QCO (value semantics): Consumes input qubit, returns reset output qubit +/// - QC (reference semantics): Resets qubit in-place, no result value +/// +/// The OpAdaptor provides the input qubit already converted to !qc.qubit. +/// Since QC's reset is in-place, we return the same qubit reference. +/// MLIR's conversion infrastructure automatically routes subsequent uses of +/// the QCO output qubit to this QC reference. +/// +/// Example transformation: +/// ```mlir +/// %q_out = qco.reset %q_in : !qco.qubit -> !qco.qubit +/// // becomes: +/// qc.reset %q : !qc.qubit +/// // %q_out uses are replaced with %q (the adaptor-converted input) +/// ``` struct ConvertQCOResetOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -647,21 +804,19 @@ struct ConvertQCOResetOp final : OpConversionPattern { } }; -/** - * @brief Converts a zero-target, one-parameter QCO gate to QC - * - * @tparam QCOOpType The operation type of the QCO gate - * @tparam QCOpType The operation type of the QC gate - * - * @par Example: - * ```mlir - * qco.gphase(%theta) - * ``` - * is converted to - * ```mlir - * qc.gphase(%theta) - * ``` - */ +/// Converts a zero-target, one-parameter QCO gate to QC +/// +/// @tparam QCOOpType The operation type of the QCO gate +/// @tparam QCOpType The operation type of the QC gate +/// +/// @par Example: +/// ```mlir +/// qco.gphase(%theta) +/// ``` +/// is converted to +/// ```mlir +/// qc.gphase(%theta) +/// ``` template struct ConvertQCOZeroTargetOneParameterToQC final : OpConversionPattern { @@ -676,19 +831,17 @@ struct ConvertQCOZeroTargetOneParameterToQC final } }; -/** - * @brief Converts qco.barrier to qc.barrier - * - * @par Example: - * ```mlir - * %q_out:2 = qco.barrier %q0_in, %q1_in : !qco.qubit, !qco.qubit -> !qco.qubit, - * !qco.qubit - * ``` - * is converted to - * ```mlir - * qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit - * ``` - */ +/// Converts qco.barrier to qc.barrier +/// +/// @par Example: +/// ```mlir +/// %q_out:2 = qco.barrier %q0_in, %q1_in : !qco.qubit, !qco.qubit -> +/// !qco.qubit, !qco.qubit +/// ``` +/// is converted to +/// ```mlir +/// qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit +/// ``` struct ConvertQCOBarrierOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -708,23 +861,21 @@ struct ConvertQCOBarrierOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.ctrl to qc.ctrl - * - * @par Example: - * ```mlir - * %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%a_in = %q1_in) { - * %a_res = qco.x %a_in : !qco.qubit -> !qco.qubit - * qco.yield %a_res : !qco.qubit - * } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) - * ``` - * is converted to - * ```mlir - * qc.ctrl(%q0) targets(%a0 = %q1) { - * qc.x %a0 : !qc.qubit - * } : !qc.qubit - * ``` - */ +/// Converts qco.ctrl to qc.ctrl +/// +/// @par Example: +/// ```mlir +/// %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%a_in = %q1_in) { +/// %a_res = qco.x %a_in : !qco.qubit -> !qco.qubit +/// qco.yield %a_res : !qco.qubit +/// } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) +/// ``` +/// is converted to +/// ```mlir +/// qc.ctrl(%q0) targets(%a0 = %q1) { +/// qc.x %a0 : !qc.qubit +/// } : !qc.qubit +/// ``` struct ConvertQCOCtrlOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -747,23 +898,21 @@ struct ConvertQCOCtrlOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.inv to qc.inv - * - * @par Example: - * ```mlir - * %q0_out = qco.inv (%a_in = %q0_in) { - * %a_res = qco.s %a_in : !qco.qubit -> !qco.qubit - * qco.yield %a_res : !qco.qubit - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - * is converted to - * ```mlir - * qc.inv { - * qc.s %q0 : !qc.qubit - * } - * ``` - */ +/// Converts qco.inv to qc.inv +/// +/// @par Example: +/// ```mlir +/// %q0_out = qco.inv (%a_in = %q0_in) { +/// %a_res = qco.s %a_in : !qco.qubit -> !qco.qubit +/// qco.yield %a_res : !qco.qubit +/// } : {!qco.qubit} -> {!qco.qubit} +/// ``` +/// is converted to +/// ```mlir +/// qc.inv { +/// qc.s %q0 : !qc.qubit +/// } +/// ``` struct ConvertQCOInvOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -785,23 +934,21 @@ struct ConvertQCOInvOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.pow to qc.pow - * - * @par Example: - * ```mlir - * %q0_out = qco.pow(%exponent) (%a_in = %q0_in) { - * %a_res = qco.s %a_in : !qco.qubit -> !qco.qubit - * qco.yield %a_res - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - * is converted to - * ```mlir - * qc.pow(%exponent) (%a0 = %q0) { - * qc.s %a0 : !qc.qubit - * } : !qc.qubit - * ``` - */ +/// Converts qco.pow to qc.pow +/// +/// @par Example: +/// ```mlir +/// %q0_out = qco.pow(%exponent) (%a_in = %q0_in) { +/// %a_res = qco.s %a_in : !qco.qubit -> !qco.qubit +/// qco.yield %a_res +/// } : {!qco.qubit} -> {!qco.qubit} +/// ``` +/// is converted to +/// ```mlir +/// qc.pow(%exponent) (%a0 = %q0) { +/// qc.s %a0 : !qc.qubit +/// } : !qc.qubit +/// ``` struct ConvertQCOPowOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -824,19 +971,17 @@ struct ConvertQCOPowOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.yield to qc.yield or to scf.yield if the parent is a - * scf::IfOp or scf::IndexSwitchOp. - * - * @par Example: - * ```mlir - * qco.yield %targets : !qco.qubit - * ``` - * is converted to - * ```mlir - * qc.yield - * ``` - */ +/// Converts qco.yield to qc.yield or to scf.yield if the parent is a +/// scf::IfOp or scf::IndexSwitchOp. +/// +/// @par Example: +/// ```mlir +/// qco.yield %targets : !qco.qubit +/// ``` +/// is converted to +/// ```mlir +/// qc.yield +/// ``` struct ConvertQCOYieldOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -858,28 +1003,26 @@ struct ConvertQCOYieldOp final : OpConversionPattern { } }; -/** - * @brief Converts scf.for with value semantics to scf.for with memory - * semantics for qubit values while preserving classical loop-carried state. - * - * @par Example: - * ```mlir - * %targets_out = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = - * %qtensor) -> (tensor<3x!qco.qubit) { - * %t0, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> - * %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit - * %insert = qtensor.insert %q1 into %t1[%iv] : tensor<3x!qco.qubit> - * scf.yield %t1 : tensor<3x!qco.qubit> - * } - * ``` - * is converted to - * ```mlir - * scf.for %iv = %lb to %ub step %step { - * %q0 = qc.load %memref[%iv] : !memref<3x!qc.qubit> - * qc.h %q0 : !qc.qubit - * } - * ``` - */ +/// Converts scf.for with value semantics to scf.for with memory +/// semantics for qubit values while preserving classical loop-carried state. +/// +/// @par Example: +/// ```mlir +/// %targets_out = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = +/// %qtensor) -> (tensor<3x!qco.qubit) { +/// %t0, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> +/// %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit +/// %insert = qtensor.insert %q1 into %t1[%iv] : tensor<3x!qco.qubit> +/// scf.yield %t1 : tensor<3x!qco.qubit> +/// } +/// ``` +/// is converted to +/// ```mlir +/// scf.for %iv = %lb to %ub step %step { +/// %q0 = qc.load %memref[%iv] : !memref<3x!qc.qubit> +/// qc.h %q0 : !qc.qubit +/// } +/// ``` struct ConvertQCOSCFForOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -910,32 +1053,30 @@ struct ConvertQCOSCFForOp final : OpConversionPattern { } }; -/** - * @brief Converts scf.while with value semantics to scf.while with memory - * semantics for qubit values while preserving classical loop-carried state. - * - * @par Example: - * ```mlir - * %targets_out = scf.while (%arg0 = %q0) : (!qco.qubit) -> !qco.qubit { - * %q1, %cond = qco.measure %arg0 : !qco.qubit - * scf.condition(%cond) %q1 : !qco.qubit - * } do { - * ^bb0(%arg0: !qco.qubit): - * %q2 = qco.h %arg0 : !qco.qubit -> !qco.qubit - * scf.yield %q2 : !qco.qubit - * } - * ``` - * is converted to - * ```mlir - * scf.while : () -> () { - * %cond = qc.measure %q0 : !qc.qubit -> i1 - * scf.condition(%cond) - * } do { - * qc.h %q0 : !qc.qubit - * scf.yield - * } - * ``` - */ +/// Converts scf.while with value semantics to scf.while with memory +/// semantics for qubit values while preserving classical loop-carried state. +/// +/// @par Example: +/// ```mlir +/// %targets_out = scf.while (%arg0 = %q0) : (!qco.qubit) -> !qco.qubit { +/// %q1, %cond = qco.measure %arg0 : !qco.qubit +/// scf.condition(%cond) %q1 : !qco.qubit +/// } do { +/// ^bb0(%arg0: !qco.qubit): +/// %q2 = qco.h %arg0 : !qco.qubit -> !qco.qubit +/// scf.yield %q2 : !qco.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// scf.while : () -> () { +/// %cond = qc.measure %q0 : !qc.qubit -> i1 +/// scf.condition(%cond) +/// } do { +/// qc.h %q0 : !qc.qubit +/// scf.yield +/// } +/// ``` struct ConvertQCOSCFWhileOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -971,25 +1112,23 @@ struct ConvertQCOSCFWhileOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.if to scf.if - * - * @par Example: - * ```mlir - * %targets_out = qco.if %cond args(%arg0 = %q0) -> (!qco.qubit) { - * %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } else args(%arg0 = %q0) { - * qco.yield %arg0 : !qco.qubit - * } - * ``` - * is converted to - * ```mlir - * scf.if %cond { - * qc.h %q0 : !qc.qubit - * } - * ``` - */ +/// Converts qco.if to scf.if +/// +/// @par Example: +/// ```mlir +/// %targets_out = qco.if %cond args(%arg0 = %q0) -> (!qco.qubit) { +/// %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q1 : !qco.qubit +/// } else args(%arg0 = %q0) { +/// qco.yield %arg0 : !qco.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// scf.if %cond { +/// qc.h %q0 : !qc.qubit +/// } +/// ``` struct ConvertQCOIfOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -1033,32 +1172,30 @@ struct ConvertQCOIfOp final : OpConversionPattern { } }; -/** - * @brief Converts qco.index_switch to scf.index_switch - * - * @par Example: - * ```mlir - * %result = qco.index_switch %condition -> !qco.qubit - * case 0 args(%arg0 = %q0) { - * %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } - * default args(%arg0 = %q0) { - * %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q2 : !qco.qubit - * } - * ``` - * is converted to - * ```mlir - * scf.index_switch %condition - * case 0 { - * qc.x %q0 : !qc.qubit - * } - * default { - * qc.z %q0 : !qc.qubit - * } - * ``` - */ +/// Converts qco.index_switch to scf.index_switch +/// +/// @par Example: +/// ```mlir +/// %result = qco.index_switch %condition -> !qco.qubit +/// case 0 args(%arg0 = %q0) { +/// %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q1 : !qco.qubit +/// } +/// default args(%arg0 = %q0) { +/// %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q2 : !qco.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// scf.index_switch %condition +/// case 0 { +/// qc.x %q0 : !qc.qubit +/// } +/// default { +/// qc.z %q0 : !qc.qubit +/// } +/// ``` struct ConvertQCOIndexSwitchOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -1092,19 +1229,17 @@ struct ConvertQCOIndexSwitchOp final : OpConversionPattern { } }; -/** - * @brief Converts scf.yield with value semantics to scf.yield with memory - * semantics for qubit values while retaining classical yielded values. - * - * @par Example: - * ```mlir - * scf.yield %targets - * ``` - * is converted to - * ```mlir - * scf.yield - * ``` - */ +/// Converts scf.yield with value semantics to scf.yield with memory +/// semantics for qubit values while retaining classical yielded values. +/// +/// @par Example: +/// ```mlir +/// scf.yield %targets +/// ``` +/// is converted to +/// ```mlir +/// scf.yield +/// ``` struct ConvertQCOSCFYieldOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -1117,19 +1252,17 @@ struct ConvertQCOSCFYieldOp final : OpConversionPattern { } }; -/** - * @brief Converts scf.condition with value semantics to scf.condition with - * memory semantics for qubit values while retaining classical state - * - * @par Example: - * ```mlir - * scf.condition(%cond) %targets - * ``` - * is converted to - * ```mlir - * scf.condition(%cond) - * ``` - */ +/// Converts scf.condition with value semantics to scf.condition with +/// memory semantics for qubit values while retaining classical state +/// +/// @par Example: +/// ```mlir +/// scf.condition(%cond) %targets +/// ``` +/// is converted to +/// ```mlir +/// scf.condition(%cond) +/// ``` struct ConvertQCOSCFConditionOp final : OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -1144,33 +1277,30 @@ struct ConvertQCOSCFConditionOp final : OpConversionPattern { } }; -/** - * @brief Pass implementation for QCO-to-QC conversion - * - * @details - * This pass converts QCO dialect operations (value semantics) to - * QC dialect operations (reference semantics). The conversion is useful - * for lowering optimized SSA-form code back to a hardware-oriented - * representation suitable for backend code generation. - * - * The conversion leverages MLIR's built-in type conversion infrastructure: - * The TypeConverter handles !qco.qubit → !qc.qubit transformations, - * and the OpAdaptor automatically provides type-converted operands to each - * conversion pattern. This eliminates the need for manual state tracking. - * - * Key semantic transformation: - * - QCO operations form explicit SSA chains where each operation consumes - * inputs and produces new outputs - * - QC operations modify qubits in-place using references - * - The conversion maps each QCO SSA chain to a single QC reference, - * with MLIR's conversion framework automatically handling the plumbing - * - * The pass operates through: - * 1. Type conversion: !qco.qubit → !qc.qubit - * 2. Operation conversion: Each QCO op converted to its QC equivalent - * 3. Automatic operand mapping: OpAdaptors provide converted operands - * 4. Function/control-flow adaptation: Signatures updated to use QC types - */ +/// Pass implementation for QCO-to-QC conversion +/// +/// This pass converts QCO dialect operations (value semantics) to +/// QC dialect operations (reference semantics). The conversion is useful +/// for lowering optimized SSA-form code back to a hardware-oriented +/// representation suitable for backend code generation. +/// +/// The conversion leverages MLIR's built-in type conversion infrastructure: +/// The TypeConverter handles !qco.qubit → !qc.qubit transformations, +/// and the OpAdaptor automatically provides type-converted operands to each +/// conversion pattern. This eliminates the need for manual state tracking. +/// +/// Key semantic transformation: +/// - QCO operations form explicit SSA chains where each operation consumes +/// inputs and produces new outputs +/// - QC operations modify qubits in-place using references +/// - The conversion maps each QCO SSA chain to a single QC reference, +/// with MLIR's conversion framework automatically handling the plumbing +/// +/// The pass operates through: +/// 1. Type conversion: !qco.qubit → !qc.qubit +/// 2. Operation conversion: Each QCO op converted to its QC equivalent +/// 3. Automatic operand mapping: OpAdaptors provide converted operands +/// 4. Function/control-flow adaptation: Signatures updated to use QC types struct QCOToQC final : impl::QCOToQCBase { using QCOToQCBase::QCOToQCBase; @@ -1181,6 +1311,23 @@ struct QCOToQC final : impl::QCOToQCBase { // Create state object to track the qubit addressing mode LoweringState state; + if (failed(collectFunctionQubitArguments(moduleOp, state))) { + signalPassFailure(); + return; + } + + SmallVector unitaryFunctions; + for (auto function : moduleOp.getOps()) { + if (mqt::isUnitaryFunction(function)) { + unitaryFunctions.emplace_back(function); + function->removeAttr(mqt::MQTDialect::UnitaryAttrHelper::getNameStr()); + } + } + auto unitaryGuard = llvm::make_scope_exit([&] { + for (auto function : unitaryFunctions) { + mqt::setUnitaryFunction(function); + } + }); ConversionTarget target(*context); RewritePatternSet patterns(context); @@ -1192,7 +1339,8 @@ struct QCOToQC final : impl::QCOToQCBase { .addLegalDialect(); target.addDynamicallyLegalDialect([](Operation* op) { - // Some types are not converted yet so QC and QCO types have to be checked + // Some types are not converted yet so QC and QCO types have to be + // checked. auto isQubitType = [](Type t) { return TypeSwitch(t) .Case([](auto) { return true; }) @@ -1231,31 +1379,30 @@ struct QCOToQC final : impl::QCOToQCBase { ConvertQTensorAllocOp, ConvertQCOAllocOp, ConvertQCOStaticOp, ConvertQCOSinkOp>(typeConverter, context, &state); - // Conversion of qco types in func.func signatures - // Note: This currently has limitations with signature changes - populateFunctionOpInterfaceTypeConversionPattern( - patterns, typeConverter); + // QCO qubit arguments are returned positionally and become in-place QC + // references again. + patterns.add(typeConverter, context, &state); target.addDynamicallyLegalOp([&](func::FuncOp op) { return typeConverter.isSignatureLegal(op.getFunctionType()) && typeConverter.isLegal(&op.getBody()); }); - // Conversion of qco types in func.return - populateReturnOpTypeConversionPattern(patterns, typeConverter); + patterns.add(typeConverter, context, &state); target.addDynamicallyLegalOp( [&](func::ReturnOp op) { return typeConverter.isLegal(op); }); - // Conversion of qco types in func.call - populateCallOpTypeConversionPattern(patterns, typeConverter); + patterns.add(typeConverter, context, &state); target.addDynamicallyLegalOp( [&](func::CallOp op) { return typeConverter.isLegal(op); }); + patterns.add(typeConverter, context); + // Conversion of qco types in control-flow ops (e.g., cf.br, cf.cond_br) populateBranchOpInterfaceTypeConversionPattern(patterns, typeConverter); - // Apply the conversion if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) { signalPassFailure(); + return; } } }; diff --git a/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp b/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp index f81b5c8a87..d2cef3316b 100644 --- a/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp +++ b/mlir/lib/Conversion/QCToQCO/QCToQCO.cpp @@ -13,6 +13,7 @@ #include "mlir/Conversion/ConversionUtils.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" #include "mlir/Dialect/CBit/IR/CBitOps.h" +#include "mlir/Dialect/MQT/IR/MQTDialect.h" #include "mlir/Dialect/QC/IR/QCDialect.h" #include "mlir/Dialect/QC/IR/QCInterfaces.h" #include "mlir/Dialect/QC/IR/QCOps.h" @@ -23,6 +24,7 @@ #include #include +#include #include #include #include @@ -66,9 +68,7 @@ namespace { using RegisterId = std::size_t; -/** - * @brief Provenance for a register-backed QC qubit reference - */ +/// Provenance for a register-backed QC qubit reference struct RegisterAccess { /// Stable identifier of the register the qubit belongs to RegisterId reg; @@ -76,48 +76,47 @@ struct RegisterAccess { Value index; }; -/** @brief Indices already used for one register by a quantum operation. */ +/// Indices already used for one register by a quantum operation. struct SeenRegisterIndices { DenseMap constants; llvm::SmallDenseSet dynamicValues; }; -/** @brief Qubit allocation mode */ +/// Qubit allocation mode enum class AllocationMode : std::uint8_t { Unset, //!< No allocation mode has been established yet. Static, //!< The module uses static qubit allocation. Dynamic //!< The module uses dynamic qubit allocation. }; -/** - * @brief State object for tracking qubit value flow during conversion - * - * @details - * This struct maintains the mapping between QC dialect qubits (which use - * reference semantics) and their corresponding QCO dialect qubit values - * (which use value semantics). As the conversion progresses, each QC - * qubit reference is mapped to its latest QCO SSA value. - * - * The key insight is that QC operations modify qubits in-place: - * ```mlir - * %q = qc.alloc : !qc.qubit - * qc.h %q : !qc.qubit // modifies %q in-place - * qc.x %q : !qc.qubit // modifies %q in-place - * ``` - * - * While QCO operations consume inputs and produce new outputs: - * ```mlir - * %q0 = qco.alloc : !qco.qubit - * %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit // %q0 consumed, %q1 produced - * %q2 = qco.x %q1 : !qco.qubit -> !qco.qubit // %q1 consumed, %q2 produced - * ``` - * - * The qubitMap tracks that the QC qubit %q corresponds to: - * - %q0 after allocation - * - %q1 after the H gate - * - %q2 after the X gate - */ +/// State object for tracking qubit value flow during conversion +/// +/// This struct maintains the mapping between QC dialect qubits (which use +/// reference semantics) and their corresponding QCO dialect qubit values +/// (which use value semantics). As the conversion progresses, each QC +/// qubit reference is mapped to its latest QCO SSA value. +/// +/// The key insight is that QC operations modify qubits in-place: +/// ```mlir +/// %q = qc.alloc : !qc.qubit +/// qc.h %q : !qc.qubit // modifies %q in-place +/// qc.x %q : !qc.qubit // modifies %q in-place +/// ``` +/// +/// While QCO operations consume inputs and produce new outputs: +/// ```mlir +/// %q0 = qco.alloc : !qco.qubit +/// %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit // %q0 consumed, %q1 produced +/// %q2 = qco.x %q1 : !qco.qubit -> !qco.qubit // %q1 consumed, %q2 produced +/// ``` +/// +/// The qubitMap tracks that the QC qubit %q corresponds to: +/// - %q0 after allocation +/// - %q1 after the H gate +/// - %q2 after the X gate struct LoweringState { + /// Original scalar-qubit arguments, retained while signatures are rewritten. + DenseMap> functionQubitArguments; struct StructuredValues { SmallVector qubits; SmallVector registers; @@ -126,7 +125,7 @@ struct LoweringState { /// Per-region map from original QC qubit reference to its latest QCO SSA /// value. /// - /// @details Keys are `Operation::getParentRegion()` for ops being converted + /// Keys are `Operation::getParentRegion()` for ops being converted /// (typically a `func.func` body or a modifier region). DenseMap> qubitMap; @@ -176,21 +175,18 @@ struct LoweringState { } }; -/** - * @brief Base class for conversion patterns that need access to lowering state - * - * @details - * Extends OpConversionPattern to provide access to a shared LoweringState - * object, which tracks the mapping from reference-semantics QC qubits - * to value-semantics QCO qubits across multiple pattern applications. - * - * This stateful approach is necessary because the conversion needs to: - * 1. Track which QCO value corresponds to each QC qubit reference - * 2. Update these mappings as operations transform qubits - * 3. Share this information across different conversion patterns - * - * @tparam OpType The QC operation type to convert - */ +/// Base class for conversion patterns that need access to lowering state +/// +/// Extends OpConversionPattern to provide access to a shared LoweringState +/// object, which tracks the mapping from reference-semantics QC qubits +/// to value-semantics QCO qubits across multiple pattern applications. +/// +/// This stateful approach is necessary because the conversion needs to: +/// 1. Track which QCO value corresponds to each QC qubit reference +/// 2. Update these mappings as operations transform qubits +/// 3. Share this information across different conversion patterns +/// +/// @tparam OpType The QC operation type to convert template class StatefulOpConversionPattern : public OpConversionPattern { @@ -207,13 +203,13 @@ class StatefulOpConversionPattern : public OpConversionPattern { }; } // namespace -/** @brief Returns whether a type is ranked or unranked QC qubit storage. */ +/// Returns whether a type is ranked or unranked QC qubit storage. [[nodiscard]] static bool isQubitMemrefType(const Type type) { const auto memref = dyn_cast(type); return memref && isa(memref.getElementType()); } -/** @brief Resolves the stable identifier for a source QC register value. */ +/// Resolves the stable identifier for a source QC register value. [[nodiscard]] static RegisterId lookupRegisterId(const LoweringState& state, Value memref) { const auto it = state.registerIds.find(memref); @@ -221,11 +217,9 @@ class StatefulOpConversionPattern : public OpConversionPattern { return it->second; } -/** - * @brief Finds the nearest region-local map containing @p reference and - * returns the pair containing the map and a mutable reference to the value in - * the map. - */ +/// Finds the nearest region-local map containing @p reference and +/// returns the pair containing the map and a mutable reference to the value in +/// the map. template [[nodiscard]] static std::pair*, Value*> findRegionLocalMap(DenseMap>& map, @@ -244,7 +238,7 @@ findRegionLocalMap(DenseMap>& map, return {nullptr, nullptr}; } -/** @brief Canonicalizes a source qubit key after block signature conversion. */ +/// Canonicalizes a source qubit key after block signature conversion. [[nodiscard]] static Value canonicalQubitKey(const LoweringState& state, Value qubit) { for (auto alias = state.convertedQubitAliases.find(qubit); @@ -255,7 +249,7 @@ findRegionLocalMap(DenseMap>& map, return qubit; } -/** @brief Resolves the latest QCO SSA value for a QC qubit reference. */ +/// Resolves the latest QCO SSA value for a QC qubit reference. [[nodiscard]] static Value lookupMappedQubit(LoweringState& state, Operation* anchor, Value qcQubit) { qcQubit = canonicalQubitKey(state, qcQubit); @@ -265,7 +259,7 @@ findRegionLocalMap(DenseMap>& map, return *qubitValue; } -/** @brief Resolves the latest QTensor SSA value for a QC register. */ +/// Resolves the latest QTensor SSA value for a QC register. [[nodiscard]] static Value lookupMappedTensor(LoweringState& state, Operation* anchor, const RegisterId reg) { @@ -276,7 +270,7 @@ findRegionLocalMap(DenseMap>& map, return *tensorValue; } -/** @brief Updates the latest QCO SSA value for a QC qubit reference. */ +/// Updates the latest QCO SSA value for a QC qubit reference. static void assignMappedQubit(LoweringState& state, Operation* anchor, Value qcQubit, Value qcoQubit) { qcQubit = canonicalQubitKey(state, qcQubit); @@ -294,7 +288,7 @@ static void assignMappedQubit(LoweringState& state, Operation* anchor, state.qubitMap[anchor->getParentRegion()][qcQubit] = qcoQubit; } -/** @brief Updates the latest QTensor SSA value for a QC register. */ +/// Updates the latest QTensor SSA value for a QC register. static void assignMappedTensor(LoweringState& state, Operation* anchor, const RegisterId reg, Value tensor) { auto [tensorMap, tensorValue] = @@ -311,7 +305,7 @@ static void assignMappedTensor(LoweringState& state, Operation* anchor, state.tensorMap[anchor->getParentRegion()][reg] = tensor; } -/** @brief Resolves a range of QC qubits to their latest QCO values. */ +/// Resolves a range of QC qubits to their latest QCO values. template [[nodiscard]] static SmallVector resolveMappedQubits(LoweringState& state, Operation* anchor, @@ -321,7 +315,7 @@ resolveMappedQubits(LoweringState& state, Operation* anchor, })); } -/** @brief Resolves a range of QC memrefs to their latest QTensor values. */ +/// Resolves a range of QC memrefs to their latest QTensor values. template [[nodiscard]] static SmallVector resolveMappedTensors(LoweringState& state, Operation* anchor, @@ -331,7 +325,7 @@ resolveMappedTensors(LoweringState& state, Operation* anchor, })); } -/** @brief Updates mappings for matching QC and QCO qubit ranges. */ +/// Updates mappings for matching QC and QCO qubit ranges. template static void assignMappedQubits(LoweringState& state, Operation* anchor, const QcRange& qcQubits, QcoRange qcoQubits) { @@ -340,7 +334,7 @@ static void assignMappedQubits(LoweringState& state, Operation* anchor, } } -/** @brief Updates mappings for matching QC memref and QTensor ranges. */ +/// Updates mappings for matching QC memref and QTensor ranges. template static void assignMappedTensors(LoweringState& state, Operation* anchor, const QcRange& registers, QcoRange tensors) { @@ -349,8 +343,7 @@ static void assignMappedTensors(LoweringState& state, Operation* anchor, } } -/** @brief Returns the structured parent whose quantum values a terminator - * yields. */ +/// Returns the structured parent whose quantum values a terminator yields. [[nodiscard]] static Operation* structuredValueOwner(Operation* operation) { if (isa(operation)) { return operation->getParentOp(); @@ -358,8 +351,7 @@ static void assignMappedTensors(LoweringState& state, Operation* anchor, return operation; } -/** @brief Seeds region-local QCO mappings for structured-control-flow block - * arguments. */ +/// Seeds region-local QCO mappings for structured-control-flow block arguments. static void seedRegionMappings(LoweringState& state, Region& region, ValueRange qcQubits, ArrayRef registers, @@ -374,7 +366,7 @@ static void seedRegionMappings(LoweringState& state, Region& region, } } -/** @brief QCO operands and register provenance materialized for a QC op. */ +/// QCO operands and register provenance materialized for a QC op. namespace { struct MaterializedQubits { SmallVector values; @@ -382,9 +374,7 @@ struct MaterializedQubits { }; } // namespace -/** - * @brief Materializes register-backed qubits immediately before a quantum op. - */ +/// Materializes register-backed qubits immediately before a quantum op. [[nodiscard]] static MaterializedQubits materializeQubits(LoweringState& state, Operation* anchor, ValueRange qcQubits, PatternRewriter& rewriter) { @@ -412,9 +402,7 @@ materializeQubits(LoweringState& state, Operation* anchor, ValueRange qcQubits, return materialized; } -/** - * @brief Commits quantum-operation results to standalone mappings or QTensor. - */ +/// Commits quantum-operation results to standalone mappings or QTensor. static void commitQubits(LoweringState& state, Operation* anchor, ValueRange qcQubits, ValueRange qcoQubits, const MaterializedQubits& materialized, @@ -437,7 +425,7 @@ static void commitQubits(LoweringState& state, Operation* anchor, } } -/** @brief Resolves all structured QC state to QCO and QTensor values. */ +/// Resolves all structured QC state to QCO and QTensor values. [[nodiscard]] static SmallVector resolveAllValues(LoweringState& state, Operation* anchor) { SmallVector registers; @@ -473,17 +461,42 @@ static void commitQubits(LoweringState& state, Operation* anchor, return success(); } -/** @brief Rejects quantum SSA sources unsupported by the lowering state. */ +/// Rejects quantum SSA sources unsupported by the lowering state. [[nodiscard]] static LogicalResult validateQuantumValueSources(Operation* root) { const auto result = root->walk([&](Operation* operation) { + if (auto returnOp = dyn_cast(operation)) { + auto function = returnOp->getParentOfType(); + llvm::SmallDenseSet returnedQubits; + for (Value value : returnOp.getOperands()) { + if (!isa(value.getType())) { + continue; + } + if (auto argument = dyn_cast(value); + argument && argument.getOwner() == &function.getBody().front()) { + returnOp.emitOpError( + "cannot return a borrowed qubit argument explicitly; QC-to-QCO " + "returns borrowed qubits implicitly"); + return WalkResult::interrupt(); + } + if (!returnedQubits.insert(value).second) { + returnOp.emitOpError("cannot return the same qubit more than once"); + return WalkResult::interrupt(); + } + } + } + const bool isModifier = isa(operation); for (Region& region : operation->getRegions()) { for (Block& block : region) { for (auto argument : block.getArguments()) { const bool isQubit = isa(argument.getType()); + const bool isFunctionArgument = + isa(operation) && + ®ion == &cast(operation).getBody() && + &block == ®ion.front(); if ((!isQubit && !isQubitMemrefType(argument.getType())) || - (isModifier && isQubit)) { + (isModifier && isQubit) || (isFunctionArgument && isQubit)) { continue; } @@ -512,7 +525,8 @@ validateQuantumValueSources(Operation* root) { } if (isa(value.getType()) && - !isa(operation)) { + !isa( + operation)) { operation->emitOpError( "produces an unsupported qubit reference; use qc.alloc, " "qc.static, a qubit-register load, or a QC modifier argument"); @@ -541,7 +555,7 @@ validateQuantumValueSources(Operation* root) { return success(!result.wasInterrupted()); } -/** @brief Collects stable register identifiers and load provenance. */ +/// Collects stable register identifiers and load provenance. [[nodiscard]] static LogicalResult collectRegisterAccesses(Operation* root, LoweringState& state) { root->walk([&](memref::AllocOp op) { @@ -572,12 +586,14 @@ collectRegisterAccesses(Operation* root, LoweringState& state) { RegisterAccess{.reg = regIt->second, .index = op.getIndices().front()}); for (Operation* user : op.getResult().getUsers()) { - if (isa(user)) { + if (isa( + user)) { continue; } user->emitOpError( - "cannot consume a register-backed qubit reference; only QC quantum " - "operations support register-backed qubits"); + "cannot consume a register-backed qubit reference; only supported " + "quantum operations and function calls support register-backed " + "qubits"); return WalkResult::interrupt(); } return WalkResult::advance(); @@ -588,14 +604,23 @@ collectRegisterAccesses(Operation* root, LoweringState& state) { } const auto distinctResult = root->walk([&](Operation* operation) { - auto unitary = dyn_cast(operation); - if (!unitary || unitary.getNumQubits() < 2) { + SmallVector operationQubits; + if (auto unitary = dyn_cast(operation)) { + llvm::append_range(operationQubits, unitary.getQubits()); + } else if (auto call = dyn_cast(operation)) { + for (auto operand : call.getOperands()) { + if (isa(operand.getType())) { + operationQubits.emplace_back(operand); + } + } + } + if (operationQubits.size() < 2) { return WalkResult::advance(); } llvm::SmallDenseSet qubits; DenseMap registerIndices; - for (auto qubit : unitary.getQubits()) { + for (auto qubit : operationQubits) { if (!qubits.insert(qubit).second) { operation->emitOpError("requires distinct qubit operands"); return WalkResult::interrupt(); @@ -633,7 +658,7 @@ collectRegisterAccesses(Operation* root, LoweringState& state) { return success(!distinctResult.wasInterrupted()); } -/** @brief Rejects unsupported operations and qubit captures in QC modifiers. */ +/// Rejects unsupported operations and qubit captures in QC modifiers. [[nodiscard]] static LogicalResult validateModifierBodies(Operation* root) { const auto result = root->walk([&](Operation* operation) { if (isa(operation)) { @@ -669,7 +694,7 @@ collectRegisterAccesses(Operation* root, LoweringState& state) { return success(!result.wasInterrupted()); } -/** @brief Collects values captured by supported structured control flow. */ +/// Collects values captured by supported structured control flow. static void collectStructuredCaptures(Operation* root, LoweringState& state) { root->walk([&](Operation* operation) { if (!isa( @@ -701,9 +726,7 @@ static void collectStructuredCaptures(Operation* root, LoweringState& state) { }); } -/** - * @brief Canonicalizes preserved SCF capture keys after signature conversion. - */ +/// Canonicalizes preserved SCF capture keys after signature conversion. static void remapStructuredCaptures(Operation* root, LoweringState& state) { root->walk([&](Operation* operation) { if (!isa( @@ -724,7 +747,7 @@ static void remapStructuredCaptures(Operation* root, LoweringState& state) { }); } -/** @brief Seeds region-owned modifier state after signature conversion. */ +/// Seeds region-owned modifier state after signature conversion. static void initializeModifierRegionState(Operation* modifier, ValueRange sourceArguments, LoweringState& state) { @@ -747,16 +770,13 @@ static void initializeModifierRegionState(Operation* modifier, namespace { -/** - * @brief Converts func.return and sinks remaining live qubits. - * - * @details - * QC uses reference semantics and does not enforce linear typing for qubits. - * After conversion, QCO requires that every qubit SSA value is consumed - * exactly once. For allocations (including static qubits), the sink is - * `qco.sink`. This pattern inserts `qco.sink` operations for all - * still-live qubits tracked in the lowering state right before the return. - */ +/// Converts func.return and sinks remaining live qubits. +/// +/// QC uses reference semantics and does not enforce linear typing for qubits. +/// After conversion, QCO requires that every qubit SSA value is consumed +/// exactly once. For allocations (including static qubits), the sink is +/// `qco.sink`. This pattern inserts `qco.sink` operations for all +/// still-live qubits tracked in the lowering state right before the return. struct ConvertFuncReturnOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -768,8 +788,8 @@ struct ConvertFuncReturnOp final : StatefulOpConversionPattern { auto& map = state.qubitMap[funcRegion]; // Build return values from qubitMap and collect live qubit information. - // A qubit from the current scope is considered alive if it is returned from - // the function. Otherwise, it is considered dead. + // A qubit from the current scope is considered alive if it is returned + // from the function. Otherwise, it is considered dead. SmallVector returnValues; returnValues.reserve(op.getNumOperands()); DenseSet liveQubits; @@ -783,6 +803,16 @@ struct ConvertFuncReturnOp final : StatefulOpConversionPattern { returnValues.emplace_back(adaptorOperand); } } + auto function = op->getParentOfType(); + for (Value argument : state.functionQubitArguments[function]) { + const auto current = map.find(argument); + if (current == map.end()) { + return op.emitOpError( + "cannot convert a function that consumes a qubit argument"); + } + returnValues.emplace_back(current->second); + liveQubits.insert(current->second); + } // Deallocate dead qubit values for (auto qcoQubit : llvm::make_second_range(map)) { @@ -797,17 +827,14 @@ struct ConvertFuncReturnOp final : StatefulOpConversionPattern { } }; -/** - * @brief Type converter for QC-to-QCO conversion - * - * @details - * Handles type conversion between the QC and QCO dialects. - * The primary conversion is from !qc.qubit to !qco.qubit, which - * represents the semantic shift from reference types to value types. - * - * Other types (integers, booleans, etc.) pass through unchanged via - * the identity conversion. - */ +/// Type converter for QC-to-QCO conversion +/// +/// Handles type conversion between the QC and QCO dialects. +/// The primary conversion is from !qc.qubit to !qco.qubit, which +/// represents the semantic shift from reference types to value types. +/// +/// Other types (integers, booleans, etc.) pass through unchanged via +/// the identity conversion. class QCToQCOTypeConverter final : public TypeConverter { public: explicit QCToQCOTypeConverter(MLIRContext* ctx) { @@ -821,18 +848,161 @@ class QCToQCOTypeConverter final : public TypeConverter { } }; -/** - * @brief Converts memref.alloc to qtensor.alloc - * - * @par Example: - * ```mlir - * %memref = memref.alloc(%c3) : memref<3x!qc.qubit> - * ``` - * is converted to - * ```mlir - * %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> - * ``` - */ +struct ConvertFuncOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(func::FuncOp op, OpAdaptor, + ConversionPatternRewriter& rewriter) const override { + if (getTypeConverter()->isSignatureLegal(op.getFunctionType())) { + return failure(); + } + + TypeConverter::SignatureConversion signature(op.getNumArguments()); + if (failed(getTypeConverter()->convertSignatureArgs(op.getArgumentTypes(), + signature))) { + return failure(); + } + SmallVector inputs; + SmallVector results; + if (failed( + getTypeConverter()->convertTypes(op.getArgumentTypes(), inputs)) || + failed( + getTypeConverter()->convertTypes(op.getResultTypes(), results))) { + return failure(); + } + + SmallVector qubitArguments; + SmallVector resultAttrs; + for (unsigned index = 0; index < op.getNumResults(); ++index) { + resultAttrs.emplace_back(op.getResultAttrDict(index)); + } + for (auto [index, type] : llvm::enumerate(op.getArgumentTypes())) { + if (isa(type)) { + qubitArguments.emplace_back(index); + results.emplace_back(qco::QubitType::get(op.getContext())); + resultAttrs.emplace_back(DictionaryAttr::get(op.getContext())); + } + } + + SmallVector originalArguments(op.getArguments()); + rewriter.modifyOpInPlace(op, [&] { + op.setType(rewriter.getFunctionType(inputs, results)); + function_interface_impl::setAllResultAttrDicts(op, resultAttrs); + }); + + if (op.isExternal()) { + return success(); + } + auto convertedEntry = rewriter.convertRegionTypes( + &op.getBody(), *getTypeConverter(), &signature); + if (failed(convertedEntry)) { + return failure(); + } + auto& map = getState().qubitMap[&op.getBody()]; + auto& functionArguments = getState().functionQubitArguments[op]; + for (unsigned index : qubitArguments) { + Value converted = (*convertedEntry)->getArgument(index); + map[originalArguments[index]] = converted; + functionArguments.emplace_back(originalArguments[index]); + } + return success(); + } +}; + +struct ConvertFuncCallOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(func::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + auto callee = SymbolTable::lookupNearestSymbolFrom( + op, op.getCalleeAttr()); + if (!callee) { + return rewriter.notifyMatchFailure(op, "callee is not defined"); + } + + auto& state = getState(); + SmallVector qcQubits; + for (auto operand : op.getOperands()) { + if (isa(operand.getType())) { + qcQubits.emplace_back(operand); + } + } + SmallVector resultTypes; + if (failed(getTypeConverter()->convertTypes(op.getResultTypes(), + resultTypes))) { + return failure(); + } + resultTypes.append(qcQubits.size(), qco::QubitType::get(op.getContext())); + + auto materialized = materializeQubits(state, op, qcQubits, rewriter); + SmallVector operands(adaptor.getOperands()); + size_t qubitIndex = 0; + for (auto [index, source] : llvm::enumerate(op.getOperands())) { + if (isa(source.getType())) { + operands[index] = materialized.values[qubitIndex++]; + } + } + auto call = func::CallOp::create(rewriter, op.getLoc(), op.getCallee(), + resultTypes, operands); + call->setAttrs(op->getAttrs()); + if (auto attrs = op.getResAttrsAttr()) { + SmallVector resultAttrs(attrs.getValue()); + resultAttrs.append(qcQubits.size(), rewriter.getDictionaryAttr({})); + call.setResAttrsAttr(rewriter.getArrayAttr(resultAttrs)); + } + + for (auto [index, source] : llvm::enumerate(op.getResults())) { + if (isa(source.getType())) { + assignMappedQubit(state, call, source, call.getResult(index)); + } + } + commitQubits(state, op, qcQubits, + call.getResults().drop_front(op.getNumResults()), materialized, + rewriter); + rewriter.replaceOp(op, call.getResults().take_front(op.getNumResults())); + return success(); + } +}; + +struct ConvertQCCallOp final : StatefulOpConversionPattern { + using StatefulOpConversionPattern::StatefulOpConversionPattern; + + LogicalResult + matchAndRewrite(qc::CallOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + auto& state = getState(); + const auto firstQubit = llvm::find_if(op.getOperands(), [](Value operand) { + return isa(operand.getType()); + }); + const auto numParams = static_cast( + std::distance(op.getOperands().begin(), firstQubit)); + auto qcQubits = op.getOperands().drop_front(numParams); + auto materialized = materializeQubits(state, op, qcQubits, rewriter); + SmallVector operands(adaptor.getOperands().take_front(numParams)); + llvm::append_range(operands, materialized.values); + auto call = qco::CallOp::create(rewriter, op.getLoc(), op.getCalleeAttr(), + operands); + call->setAttrs(op->getAttrs()); + call.removeResAttrsAttr(); + commitQubits(state, op, qcQubits, call.getOutputQubits(), materialized, + rewriter); + rewriter.eraseOp(op); + return success(); + } +}; + +/// Converts memref.alloc to qtensor.alloc +/// +/// @par Example: +/// ```mlir +/// %memref = memref.alloc(%c3) : memref<3x!qc.qubit> +/// ``` +/// is converted to +/// ```mlir +/// %tensor = qtensor.alloc(%c3) : tensor<3x!qco.qubit> +/// ``` struct ConvertMemRefAllocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -874,16 +1044,14 @@ struct ConvertMemRefAllocOp final } }; -/** - * @brief Erases a qubit memref.load after recording its converted index - * - * @par Example: - * ```mlir - * %q = memref.load %memref[%c0] : memref<3x!qc.qubit> - * ``` - * The consuming quantum operation materializes and commits the referenced - * qubit locally. - */ +/// Erases a qubit memref.load after recording its converted index +/// +/// @par Example: +/// ```mlir +/// %q = memref.load %memref[%c0] : memref<3x!qc.qubit> +/// ``` +/// The consuming quantum operation materializes and commits the referenced +/// qubit locally. struct ConvertMemRefLoadOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -906,18 +1074,16 @@ struct ConvertMemRefLoadOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts memref.dealloc to qtensor.dealloc - * - * @par Example: - * ```mlir - * memref.dealloc %memref : memref<3x!qc.qubit> - * ``` - * is converted to - * ```mlir - * qtensor.dealloc %tensor : tensor<3x!qco.qubit> - * ``` - */ +/// Converts memref.dealloc to qtensor.dealloc +/// +/// @par Example: +/// ```mlir +/// memref.dealloc %memref : memref<3x!qc.qubit> +/// ``` +/// is converted to +/// ```mlir +/// qtensor.dealloc %tensor : tensor<3x!qco.qubit> +/// ``` struct ConvertMemRefDeallocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -941,18 +1107,16 @@ struct ConvertMemRefDeallocOp final } }; -/** - * @brief Converts qc.alloc to qco.alloc - * - * @par Example: - * ```mlir - * %q = qc.alloc : !qc.qubit - * ``` - * is converted to - * ```mlir - * %q = qco.alloc : !qco.qubit - * ``` - */ +/// Converts qc.alloc to qco.alloc +/// +/// @par Example: +/// ```mlir +/// %q = qc.alloc : !qc.qubit +/// ``` +/// is converted to +/// ```mlir +/// %q = qco.alloc : !qco.qubit +/// ``` struct ConvertQCAllocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -976,21 +1140,18 @@ struct ConvertQCAllocOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.dealloc to qco.sink - * - * @details - * Deallocates a qubit by looking up its latest QCO value and creating - * a corresponding qco.sink operation. The mapping is removed from - * the state as the qubit is no longer in use. - * - * Example transformation: - * ```mlir - * qc.dealloc %q : !qc.qubit - * // becomes (where %q maps to %q_final): - * qco.sink %q_final : !qco.qubit - * ``` - */ +/// Converts qc.dealloc to qco.sink +/// +/// Deallocates a qubit by looking up its latest QCO value and creating +/// a corresponding qco.sink operation. The mapping is removed from +/// the state as the qubit is no longer in use. +/// +/// Example transformation: +/// ```mlir +/// qc.dealloc %q : !qc.qubit +/// // becomes (where %q maps to %q_final): +/// qco.sink %q_final : !qco.qubit +/// ``` struct ConvertQCDeallocOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1014,21 +1175,18 @@ struct ConvertQCDeallocOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.static to qco.static - * - * @details - * Static qubits represent references to hardware-mapped or fixed-position - * qubits identified by an index. This conversion creates the corresponding - * qco.static operation and establishes the mapping. - * - * Example transformation: - * ```mlir - * %q = qc.static 0 : !qc.qubit - * // becomes: - * %q0 = qco.static 0 : !qco.qubit - * ``` - */ +/// Converts qc.static to qco.static +/// +/// Static qubits represent references to hardware-mapped or fixed-position +/// qubits identified by an index. This conversion creates the corresponding +/// qco.static operation and establishes the mapping. +/// +/// Example transformation: +/// ```mlir +/// %q = qc.static 0 : !qc.qubit +/// // becomes: +/// %q0 = qco.static 0 : !qco.qubit +/// ``` struct ConvertQCStaticOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1049,27 +1207,24 @@ struct ConvertQCStaticOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.measure to qco.measure - * - * @details - * Measurement is a key operation where the semantic difference is visible: - * - QC: Measures in-place, returning only the classical bit - * - QCO: Consumes input qubit, returns both output qubit and classical bit - * - * The conversion looks up the latest QCO value for the QC qubit, - * performs the measurement, updates the mapping with the output qubit, - * and returns the classical bit result. - * - * @par Example: - * ```mlir - * %c = qc.measure %q : !qc.qubit -> i1 - * ``` - * is converted to - * ```mlir - * %q_out, %c = qco.measure %q_in : !qco.qubit - * ``` - */ +/// Converts qc.measure to qco.measure +/// +/// Measurement is a key operation where the semantic difference is visible: +/// - QC: Measures in-place, returning only the classical bit +/// - QCO: Consumes input qubit, returns both output qubit and classical bit +/// +/// The conversion looks up the latest QCO value for the QC qubit, +/// performs the measurement, updates the mapping with the output qubit, +/// and returns the classical bit result. +/// +/// @par Example: +/// ```mlir +/// %c = qc.measure %q : !qc.qubit -> i1 +/// ``` +/// is converted to +/// ```mlir +/// %q_out, %c = qco.measure %q_in : !qco.qubit +/// ``` struct ConvertQCMeasureOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1096,26 +1251,23 @@ struct ConvertQCMeasureOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.reset to qco.reset - * - * @details - * Reset operations force a qubit to the |0⟩ state. The semantic difference: - * - QC: Resets in-place (no result value) - * - QCO: Consumes input qubit, returns reset output qubit - * - * The conversion looks up the latest QCO value, performs the reset, - * and updates the mapping with the output qubit. The QC operation - * is erased as it has no results to replace. - * - * Example transformation: - * ```mlir - * qc.reset %q : !qc.qubit - * // becomes (where %q maps to %q_in): - * %q_out = qco.reset %q_in : !qco.qubit -> !qco.qubit - * // state updated: %q now maps to %q_out - * ``` - */ +/// Converts qc.reset to qco.reset +/// +/// Reset operations force a qubit to the |0⟩ state. The semantic difference: +/// - QC: Resets in-place (no result value) +/// - QCO: Consumes input qubit, returns reset output qubit +/// +/// The conversion looks up the latest QCO value, performs the reset, +/// and updates the mapping with the output qubit. The QC operation +/// is erased as it has no results to replace. +/// +/// Example transformation: +/// ```mlir +/// qc.reset %q : !qc.qubit +/// // becomes (where %q maps to %q_in): +/// %q_out = qco.reset %q_in : !qco.qubit -> !qco.qubit +/// // state updated: %q now maps to %q_out +/// ``` struct ConvertQCResetOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1177,7 +1329,7 @@ struct ConvertQCGateToQCO final : StatefulOpConversionPattern { } }; -/** Converts a variadic dense qc.unitary to its value-semantics form. */ +/// Converts a variadic dense qc.unitary to its value-semantics form. struct ConvertQCUnitaryOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1198,19 +1350,17 @@ struct ConvertQCUnitaryOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.barrier to qco.barrier - * - * @par Example: - * ```mlir - * qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit - * ``` - * is converted to - * ```mlir - * %q_out:2 = qco.barrier %q0_in, %q1_in : !qco.qubit, !qco.qubit -> !qco.qubit, - * !qco.qubit - * ``` - */ +/// Converts qc.barrier to qco.barrier +/// +/// @par Example: +/// ```mlir +/// qc.barrier %q0, %q1 : !qc.qubit, !qc.qubit +/// ``` +/// is converted to +/// ```mlir +/// %q_out:2 = qco.barrier %q0_in, %q1_in : !qco.qubit, !qco.qubit -> +/// !qco.qubit, !qco.qubit +/// ``` struct ConvertQCBarrierOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1234,23 +1384,21 @@ struct ConvertQCBarrierOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.ctrl to qco.ctrl - * - * @par Example: - * ```mlir - * qc.ctrl(%q0) targets(%a0 = %q1) { - * qc.x %a0 : !qc.qubit - * } : !qc.qubit - * ``` - * is converted to - * ```mlir - * %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%a_in = %q1_in) { - * %a_res = qco.x %a_in : !qco.qubit -> !qco.qubit - * qco.yield %a_res : !qco.qubit - * } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) - * ``` - */ +/// Converts qc.ctrl to qco.ctrl +/// +/// @par Example: +/// ```mlir +/// qc.ctrl(%q0) targets(%a0 = %q1) { +/// qc.x %a0 : !qc.qubit +/// } : !qc.qubit +/// ``` +/// is converted to +/// ```mlir +/// %controls_out, %targets_out = qco.ctrl(%q0_in) targets(%a_in = %q1_in) { +/// %a_res = qco.x %a_in : !qco.qubit -> !qco.qubit +/// qco.yield %a_res : !qco.qubit +/// } : ({!qco.qubit}, {!qco.qubit}) -> ({!qco.qubit}, {!qco.qubit}) +/// ``` struct ConvertQCCtrlOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1292,23 +1440,21 @@ struct ConvertQCCtrlOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.inv to qco.inv - * - * @par Example: - * ```mlir - * qc.inv { - * qc.s %q0 : !qc.qubit - * } : !qc.qubit - * ``` - * is converted to - * ```mlir - * %q0_out = qco.inv (%a0_in = %q0_in) { - * %a0_res = qco.s %a0_in : !qco.qubit -> !qco.qubit - * qco.yield %a0_res : !qco.qubit - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - */ +/// Converts qc.inv to qco.inv +/// +/// @par Example: +/// ```mlir +/// qc.inv { +/// qc.s %q0 : !qc.qubit +/// } : !qc.qubit +/// ``` +/// is converted to +/// ```mlir +/// %q0_out = qco.inv (%a0_in = %q0_in) { +/// %a0_res = qco.s %a0_in : !qco.qubit -> !qco.qubit +/// qco.yield %a0_res : !qco.qubit +/// } : {!qco.qubit} -> {!qco.qubit} +/// ``` struct ConvertQCInvOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1343,23 +1489,21 @@ struct ConvertQCInvOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.pow to qco.pow - * - * @par Example: - * ```mlir - * qc.pow(%exponent) (%a0 = %q0) { - * qc.s %a0 : !qc.qubit - * } : !qc.qubit - * ``` - * is converted to - * ```mlir - * %q0_out = qco.pow(%exponent) (%a0_in = %q0_in) { - * %a0_res = qco.s %a0_in : !qco.qubit -> !qco.qubit - * qco.yield %a0_res - * } : {!qco.qubit} -> {!qco.qubit} - * ``` - */ +/// Converts qc.pow to qco.pow +/// +/// @par Example: +/// ```mlir +/// qc.pow(%exponent) (%a0 = %q0) { +/// qc.s %a0 : !qc.qubit +/// } : !qc.qubit +/// ``` +/// is converted to +/// ```mlir +/// %q0_out = qco.pow(%exponent) (%a0_in = %q0_in) { +/// %a0_res = qco.s %a0_in : !qco.qubit -> !qco.qubit +/// qco.yield %a0_res +/// } : {!qco.qubit} -> {!qco.qubit} +/// ``` struct ConvertQCPowOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1395,18 +1539,16 @@ struct ConvertQCPowOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts qc.yield to qco.yield - * - * @par Example: - * ```mlir - * qc.yield - * ``` - * is converted to - * ```mlir - * qco.yield %targets : !qco.qubit - * ``` - */ +/// Converts qc.yield to qco.yield +/// +/// @par Example: +/// ```mlir +/// qc.yield +/// ``` +/// is converted to +/// ```mlir +/// qco.yield %targets : !qco.qubit +/// ``` struct ConvertQCYieldOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1429,28 +1571,26 @@ struct ConvertQCYieldOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts scf.for with memory semantics to scf.for with value - * semantics for qubit values - * - * @par Example: - * ```mlir - * scf.for %iv = %lb to %ub step %step { - * %q0 = qc.load %memref[%iv] : !memref<3x!qc.qubit> - * qc.h %q0 : !qc.qubit - * } - * ``` - * is converted to - * ```mlir - * %targets_out = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = - * %qtensor) -> (tensor<3x!qco.qubit) { - * %t0, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> - * %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit - * %t1 = qtensor.insert %q1 into %t0[%iv] : tensor<3x!qco.qubit> - * scf.yield %t1 : tensor<3x!qco.qubit> - * } - * ``` - */ +/// Converts scf.for with memory semantics to scf.for with value +/// semantics for qubit values +/// +/// @par Example: +/// ```mlir +/// scf.for %iv = %lb to %ub step %step { +/// %q0 = qc.load %memref[%iv] : !memref<3x!qc.qubit> +/// qc.h %q0 : !qc.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// %targets_out = scf.for %iv = %lb to %ub step %step iter_args(%arg0 = +/// %qtensor) -> (tensor<3x!qco.qubit) { +/// %t0, %q0 = qtensor.extract %arg0[%iv] : tensor<3x!qco.qubit> +/// %q1 = qco.h %q0 : !qco.qubit -> !qco.qubit +/// %t1 = qtensor.insert %q1 into %t0[%iv] : tensor<3x!qco.qubit> +/// scf.yield %t1 : tensor<3x!qco.qubit> +/// } +/// ``` struct ConvertSCFForOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1509,32 +1649,30 @@ struct ConvertSCFForOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts scf.while with memory semantics to scf.while with value - * semantics for qubit values. - * - * @par Example: - * ```mlir - * scf.while : () -> () { - * %cond = qc.measure %q0 : !qc.qubit -> i1 - * scf.condition(%cond) - * } do { - * qc.h %q0 : !qc.qubit - * scf.yield - * } - * ``` - * is converted to - * ```mlir - * %targets_out = scf.while (%arg0 = %q0) : (!qco.qubit) -> !qco.qubit { - * %q1 = qco.measure %arg0 : !qco.qubit - * scf.condition(%cond) %q1 : !qco.qubit - * } do { - * ^bb0(%arg0: !qco.qubit): - * %q2 = qco.h %arg0 : !qco.qubit -> !qco.qubit - * scf.yield %q2 : !qco.qubit - * } - * ``` - */ +/// Converts scf.while with memory semantics to scf.while with value +/// semantics for qubit values. +/// +/// @par Example: +/// ```mlir +/// scf.while : () -> () { +/// %cond = qc.measure %q0 : !qc.qubit -> i1 +/// scf.condition(%cond) +/// } do { +/// qc.h %q0 : !qc.qubit +/// scf.yield +/// } +/// ``` +/// is converted to +/// ```mlir +/// %targets_out = scf.while (%arg0 = %q0) : (!qco.qubit) -> !qco.qubit { +/// %q1 = qco.measure %arg0 : !qco.qubit +/// scf.condition(%cond) %q1 : !qco.qubit +/// } do { +/// ^bb0(%arg0: !qco.qubit): +/// %q2 = qco.h %arg0 : !qco.qubit -> !qco.qubit +/// scf.yield %q2 : !qco.qubit +/// } +/// ``` struct ConvertSCFWhileOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1621,25 +1759,23 @@ struct ConvertSCFWhileOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts scf.if to qco.if - * - * @par Example: - * ```mlir - * scf.if %cond { - * qc.h %q0 : !qc.qubit - * } - * ``` - * is converted to - * ```mlir - * %targets_out = qco.if %cond args(%arg0 = %q0) -> (!qco.qubit) { - * %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } else args(%arg0 = %q0) { - * qco.yield %arg0 : !qco.qubit - * } - * ``` - */ +/// Converts scf.if to qco.if +/// +/// @par Example: +/// ```mlir +/// scf.if %cond { +/// qc.h %q0 : !qc.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// %targets_out = qco.if %cond args(%arg0 = %q0) -> (!qco.qubit) { +/// %q1 = qco.h %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q1 : !qco.qubit +/// } else args(%arg0 = %q0) { +/// qco.yield %arg0 : !qco.qubit +/// } +/// ``` struct ConvertSCFIfOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1707,32 +1843,30 @@ struct ConvertSCFIfOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts scf.index_switch to qco.index_switch - * - * @par Example: - * ```mlir - * scf.index_switch %condition - * case 0 { - * qc.x %q0 : !qc.qubit - * } - * default { - * qc.z %q0 : !qc.qubit - * } - * ``` - * is converted to - * ```mlir - * %result = qco.index_switch %condition -> !qco.qubit - * case 0 args(%arg0 = %q0) { - * %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q1 : !qco.qubit - * } - * default args(%arg0 = %q0) { - * %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit - * qco.yield %q2 : !qco.qubit - * } - * ``` - */ +/// Converts scf.index_switch to qco.index_switch +/// +/// @par Example: +/// ```mlir +/// scf.index_switch %condition +/// case 0 { +/// qc.x %q0 : !qc.qubit +/// } +/// default { +/// qc.z %q0 : !qc.qubit +/// } +/// ``` +/// is converted to +/// ```mlir +/// %result = qco.index_switch %condition -> !qco.qubit +/// case 0 args(%arg0 = %q0) { +/// %q1 = qco.x %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q1 : !qco.qubit +/// } +/// default args(%arg0 = %q0) { +/// %q2 = qco.z %arg0 : !qco.qubit -> !qco.qubit +/// qco.yield %q2 : !qco.qubit +/// } +/// ``` struct ConvertSCFIndexSwitchOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1789,20 +1923,18 @@ struct ConvertSCFIndexSwitchOp final } }; -/** - * @brief Converts scf.yield with memory semantics to scf.yield with value - * semantics for qubit values or to qco.yield if the parentOp is a qco::IfOp or - * qco::IndexSwitchOp. - * - * @par Example: - * ```mlir - * scf.yield - * ``` - * is converted to - * ```mlir - * scf.yield %targets - * ``` - */ +/// Converts scf.yield with memory semantics to scf.yield with value +/// semantics for qubit values or to qco.yield if the parentOp is a qco::IfOp or +/// qco::IndexSwitchOp. +/// +/// @par Example: +/// ```mlir +/// scf.yield +/// ``` +/// is converted to +/// ```mlir +/// scf.yield %targets +/// ``` struct ConvertSCFYieldOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1825,19 +1957,17 @@ struct ConvertSCFYieldOp final : StatefulOpConversionPattern { } }; -/** - * @brief Converts scf.condition with memory semantics to scf.condition with - * value semantics for qubit values - * - * @par Example: - * ```mlir - * scf.condition(%cond) - * ``` - * is converted to - * ```mlir - * scf.condition(%cond) %targets - * ``` - */ +/// Converts scf.condition with memory semantics to scf.condition with +/// value semantics for qubit values +/// +/// @par Example: +/// ```mlir +/// scf.condition(%cond) +/// ``` +/// is converted to +/// ```mlir +/// scf.condition(%cond) %targets +/// ``` struct ConvertSCFConditionOp final : StatefulOpConversionPattern { using StatefulOpConversionPattern::StatefulOpConversionPattern; @@ -1857,24 +1987,22 @@ struct ConvertSCFConditionOp final } }; -/** - * @brief Pass implementation for QC-to-QCO conversion - * - * @details - * This pass converts QC dialect operations (reference semantics) to QCO dialect - * operations (value semantics). The conversion is essential for enabling - * optimization passes that rely on SSA form and explicit dataflow analysis. - * - * The pass operates in several phases: - * 1. Type conversion: !qc.qubit -> !qco.qubit - * 2. Operation conversion: Each QC op is converted to its QCO equivalent - * 3. State tracking: A LoweringState maintains qubit value mappings - * 4. Function/control-flow adaptation: Function signatures and control flow are - * updated to use QCO types - * - * The conversion maintains semantic equivalence while transforming the - * representation from imperative (mutation-based) to functional (SSA-based). - */ +/// Pass implementation for QC-to-QCO conversion +/// +/// This pass converts QC dialect operations (reference semantics) to QCO +/// dialect operations (value semantics). The conversion is essential for +/// enabling optimization passes that rely on SSA form and explicit dataflow +/// analysis. +/// +/// The pass operates in several phases: +/// 1. Type conversion: !qc.qubit -> !qco.qubit +/// 2. Operation conversion: Each QC op is converted to its QCO equivalent +/// 3. State tracking: A LoweringState maintains qubit value mappings +/// 4. Function/control-flow adaptation: Function signatures and control flow +/// are updated to use QCO types +/// +/// The conversion maintains semantic equivalence while transforming the +/// representation from imperative (mutation-based) to functional (SSA-based). struct QCToQCO final : impl::QCToQCOBase { using QCToQCOBase::QCToQCOBase; @@ -1908,6 +2036,19 @@ struct QCToQCO final : impl::QCToQCOBase { return; } + SmallVector unitaryFunctions; + for (auto function : moduleOp.getOps()) { + if (mqt::isUnitaryFunction(function)) { + unitaryFunctions.emplace_back(function); + function->removeAttr(mqt::MQTDialect::UnitaryAttrHelper::getNameStr()); + } + } + auto unitaryGuard = llvm::make_scope_exit([&] { + for (auto function : unitaryFunctions) { + mqt::setUnitaryFunction(function); + } + }); + // Get the quantum values captured by structured control-flow regions. collectStructuredCaptures(moduleOp, state); @@ -1940,8 +2081,8 @@ struct QCToQCO final : impl::QCToQCOBase { ConvertMemRefDeallocOp, ConvertQCAllocOp, ConvertQCDeallocOp, ConvertQCStaticOp, ConvertQCMeasureOp, ConvertQCResetOp, ConvertQCUnitaryOp, ConvertQCBarrierOp, ConvertQCCtrlOp, - ConvertQCInvOp, ConvertQCPowOp, ConvertQCYieldOp>(typeConverter, - context, &state); + ConvertQCInvOp, ConvertQCPowOp, ConvertQCYieldOp, ConvertQCCallOp>( + typeConverter, context, &state); // Not part of the central gate table. patterns.add>( @@ -1953,10 +2094,9 @@ struct QCToQCO final : impl::QCToQCOBase { typeConverter, context, &state); #include "mlir/Conversion/GateTable.def" - // Conversion of qc types in func.func signatures - // Note: This currently has limitations with signature changes - populateFunctionOpInterfaceTypeConversionPattern( - patterns, typeConverter); + // QC qubit arguments become QCO arguments plus trailing pass-through + // results. + patterns.add(typeConverter, context, &state); target.addDynamicallyLegalOp([&](func::FuncOp op) { return typeConverter.isSignatureLegal(op.getFunctionType()) && typeConverter.isLegal(&op.getBody()); @@ -1976,8 +2116,8 @@ struct QCToQCO final : impl::QCToQCOBase { return it == state.qubitMap.end() || it->second.empty(); }); - // Conversion of qc types in func.call - populateCallOpTypeConversionPattern(patterns, typeConverter); + // Generic calls receive the pass-through results added to their callees. + patterns.add(typeConverter, context, &state); target.addDynamicallyLegalOp( [&](func::CallOp op) { return typeConverter.isLegal(op); }); @@ -1989,7 +2129,6 @@ struct QCToQCO final : impl::QCToQCOBase { signalPassFailure(); return; } - // Source register values and loaded qubit references have been erased. // Structured conversion state uses stable register identifiers from here // on. diff --git a/mlir/lib/Dialect/MQT/IR/MQTDialect.cpp b/mlir/lib/Dialect/MQT/IR/MQTDialect.cpp index 037ee60f49..556dc5e863 100644 --- a/mlir/lib/Dialect/MQT/IR/MQTDialect.cpp +++ b/mlir/lib/Dialect/MQT/IR/MQTDialect.cpp @@ -13,7 +13,11 @@ #include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/MQT/IR/MQTAttributes.h" #include "mlir/Dialect/QC/IR/QCDialect.h" +#include "mlir/Dialect/QC/IR/QCInterfaces.h" +#include "mlir/Dialect/QC/IR/QCOps.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" +#include "mlir/Dialect/QCO/IR/QCOInterfaces.h" +#include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include @@ -21,6 +25,7 @@ #include #include #include // IWYU pragma: keep +#include #include #include #include @@ -31,7 +36,10 @@ #include #include // IWYU pragma: keep #include +#include +#include #include +#include #include #include @@ -304,6 +312,192 @@ verifyEntryPoint(Operation* operation, const NamedAttribute attribute) { return success(); } +template +[[nodiscard]] static LogicalResult +verifyNoUnitaryRecursion(func::FuncOp function) { + // A cycle of unitary calls passes local body checks without any control flow. + DenseSet visited; + SmallVector worklist{function}; + while (!worklist.empty()) { + auto current = worklist.pop_back_val(); + if (!visited.insert(current).second) { + continue; + } + WalkResult result = current.walk([&](CallOp call) { + if (failed(verify(call))) { + return WalkResult::interrupt(); + } + auto callee = SymbolTable::lookupNearestSymbolFrom( + call, call.getCalleeAttr()); + if (!callee) { + return WalkResult::advance(); + } + if (callee == function) { + function.emitError("unitary function must not be recursive"); + return WalkResult::interrupt(); + } + worklist.emplace_back(callee); + return WalkResult::advance(); + }); + if (result.wasInterrupted()) { + return failure(); + } + } + return success(); +} + +[[nodiscard]] static LogicalResult verifyQCUnitaryBody(func::FuncOp function) { + bool valid = true; + function.walk([&](Operation* nested) { + if (!valid || nested == function.getOperation()) { + return; + } + if (isa(nested)) { + return; + } + if (isa(nested)) { + return; + } + valid = + nested->getNumRegions() == 0 && isMemoryEffectFree(nested) && + llvm::none_of(nested->getOperandTypes(), + llvm::IsaPred) && + llvm::none_of(nested->getResultTypes(), llvm::IsaPred); + }); + if (!valid) { + return function.emitError() + << "unitary QC function body contains a non-unitary operation"; + } + + return verifyNoUnitaryRecursion(function); +} + +[[nodiscard]] static LogicalResult verifyQCOUnitaryBody(func::FuncOp function, + unsigned firstQubit) { + bool valid = true; + function.walk([&](Operation* nested) { + if (!valid || nested == function.getOperation()) { + return; + } + if (isa(nested)) { + return; + } + valid = + nested->getNumRegions() == 0 && isMemoryEffectFree(nested) && + llvm::none_of(nested->getOperandTypes(), + llvm::IsaPred) && + llvm::none_of(nested->getResultTypes(), llvm::IsaPred); + }); + if (!valid) { + return function.emitError() + << "unitary QCO function body contains a non-unitary operation"; + } + + auto returnOp = cast(function.getBody().front().back()); + for (auto [resultIndex, returned] : llvm::enumerate(returnOp.getOperands())) { + Value current = returned; + llvm::SmallDenseSet visited; + while (auto result = dyn_cast(current)) { + if (!visited.insert(current).second) { + return function.emitError("unitary QCO result has cyclic qubit flow"); + } + auto unitary = dyn_cast(result.getOwner()); + if (!unitary) { + return function.emitError() + << "unitary QCO result does not originate from a qubit " + "argument"; + } + current = unitary.getInputForOutput(current); + if (!current) { + return function.emitError() + << "unitary QCO operation has no input corresponding to its " + "returned qubit"; + } + } + auto argument = dyn_cast(current); + if (!argument || argument.getOwner() != &function.getBody().front() || + argument.getArgNumber() != firstQubit + resultIndex) { + return function.emitError() + << "unitary QCO results must continue qubit arguments " + "positionally"; + } + } + return verifyNoUnitaryRecursion(function); +} + +[[nodiscard]] static LogicalResult +verifyUnitaryFunction(Operation* operation, const NamedAttribute attribute) { + if (!isa(attribute.getValue())) { + return operation->emitError() + << "attribute '" << attribute.getName().getValue() + << "' must be a unit attribute"; + } + + auto function = dyn_cast(operation); + if (!function || function.isExternal() || !function.isPrivate() || + isEntryPoint(operation) || !function.getBody().hasOneBlock()) { + return operation->emitError() + << "attribute '" << attribute.getName().getValue() + << "' requires a private, defined, single-block non-entry function"; + } + if (function.getBody().front().empty()) { + return operation->emitError("unitary function body must not be empty"); + } + + unsigned firstQubit = function.getNumArguments(); + bool usesQC = false; + bool usesQCO = false; + for (auto [index, type] : llvm::enumerate(function.getArgumentTypes())) { + if (isa(type)) { + if (firstQubit == function.getNumArguments()) { + firstQubit = index; + } + usesQC |= isa(type); + usesQCO |= isa(type); + continue; + } + if (firstQubit != function.getNumArguments() || !type.isF64()) { + return operation->emitError() + << "unitary function arguments must be f64 parameters " + "followed by scalar qubits"; + } + } + if (firstQubit == function.getNumArguments() || usesQC == usesQCO) { + return operation->emitError() + << "unitary function requires at least one QC or QCO qubit " + "argument"; + } + + const auto numQubits = function.getNumArguments() - firstQubit; + if (usesQC && function.getNumResults() != 0) { + return operation->emitError() + << "unitary QC function must not return values"; + } + if (usesQCO && (function.getNumResults() != numQubits || + llvm::any_of(function.getResultTypes(), [](Type type) { + return !isa(type); + }))) { + return operation->emitError() + << "unitary QCO function must return one qubit per qubit argument"; + } + auto returnOp = dyn_cast(function.getBody().front().back()); + if (!returnOp || (usesQC && returnOp.getNumOperands() != 0)) { + return operation->emitError( + usesQC ? "unitary QC function must end in an empty func.return" + : "unitary QCO function must end in func.return"); + } + + // Attribute verification precedes nested operation verification. Check the + // body before querying memory effects or qubit correspondence. + for (Operation& nested : function.getBody().front()) { + if (failed(verify(&nested))) { + return failure(); + } + } + return usesQC ? verifyQCUnitaryBody(function) + : verifyQCOUnitaryBody(function, firstQubit); +} + [[nodiscard]] static LogicalResult verifyName(Operation* operation, const NamedAttribute attribute) { const auto name = dyn_cast(attribute.getValue()); @@ -440,9 +634,20 @@ MQTDialect::verifyOperationAttribute(Operation* operation, if (attribute.getName() == EntryPointAttrHelper::getNameStr()) { return verifyEntryPoint(operation, attribute); } + if (attribute.getName() == UnitaryAttrHelper::getNameStr()) { + return verifyUnitaryFunction(operation, attribute); + } if (attribute.getName() == RegisterNameAttrHelper::getNameStr()) { return verifyRegisterName(operation, attribute); } + if (attribute.getName() == SourceNameAttrHelper::getNameStr()) { + if (!isa(operation)) { + return operation->emitError() + << "attribute '" << attribute.getName().getValue() + << "' is only valid on a function"; + } + return verifyName(operation, attribute); + } if (attribute.getName() == ParameterGroupAttrHelper::getNameStr()) { if (!isa(operation)) { return operation->emitError() @@ -525,3 +730,8 @@ void mlir::mqt::setEntryPoint(Operation* operation) { void mlir::mqt::removeEntryPoint(Operation* operation) { operation->removeAttr(MQTDialect::EntryPointAttrHelper::getNameStr()); } + +void mlir::mqt::setUnitaryFunction(Operation* operation) { + operation->setAttr(MQTDialect::UnitaryAttrHelper::getNameStr(), + UnitAttr::get(operation->getContext())); +} diff --git a/mlir/lib/Dialect/QC/Builder/QCProgramBuilder.cpp b/mlir/lib/Dialect/QC/Builder/QCProgramBuilder.cpp index de320801de..6419f0027e 100644 --- a/mlir/lib/Dialect/QC/Builder/QCProgramBuilder.cpp +++ b/mlir/lib/Dialect/QC/Builder/QCProgramBuilder.cpp @@ -19,6 +19,7 @@ #include "mlir/Dialect/QC/IR/QCOps.h" #include +#include #include #include #include @@ -31,6 +32,7 @@ #include #include #include +#include #include #include #include @@ -47,7 +49,7 @@ namespace mlir::qc { QCProgramBuilder::QCProgramBuilder(MLIRContext* context) : ImplicitLocOpBuilder( FileLineColLoc::get(context, "", 1, 1), context), - ctx(context), module(ModuleOp::create(*this)) { + ctx(context), moduleOp_(ModuleOp::create(*this)) { ctx->loadDialect(); } @@ -55,7 +57,7 @@ void QCProgramBuilder::initialize() { initialize({getI64Type()}); } void QCProgramBuilder::initialize(TypeRange returnTypes) { // Set insertion point to the module body - setInsertionPointToStart(cast(module).getBody()); + setInsertionPointToStart(cast(moduleOp_).getBody()); // Create main function as entry point auto funcType = getFunctionType({}, returnTypes); @@ -69,7 +71,7 @@ void QCProgramBuilder::initialize(TypeRange returnTypes) { } void QCProgramBuilder::retype(TypeRange returnTypes) { - auto mainFunc = mqt::getEntryPoint(cast(module)); + auto mainFunc = mqt::getEntryPoint(cast(moduleOp_)); if (!mainFunc) { llvm::reportFatalUsageError("Main function not found for retyping"); } @@ -78,6 +80,93 @@ void QCProgramBuilder::retype(TypeRange returnTypes) { mainFunc.setType(funcType); } +func::FuncOp QCProgramBuilder::createFunction( + const StringRef name, const TypeRange argumentTypes, + const function_ref(ValueRange)> body) { + checkFinalized(); + auto moduleOp = cast(moduleOp_); + auto mainFunc = mqt::getEntryPoint(moduleOp); + if (!mainFunc) { + llvm::reportFatalUsageError( + "QCProgramBuilder must be initialized before creating a function"); + } + if (SymbolTable::lookupSymbolIn(moduleOp, name) != nullptr) { + llvm::reportFatalUsageError("Function name is already defined"); + } + + const InsertionGuard insertionGuard(*this); + auto savedAllocatedQubits = std::move(allocatedQubits); + auto savedAllocatedQregs = std::move(allocatedQregs); + auto savedStaticQubits = std::move(staticQubits); + auto stateGuard = llvm::make_scope_exit([&] { + allocatedQubits = std::move(savedAllocatedQubits); + allocatedQregs = std::move(savedAllocatedQregs); + staticQubits = std::move(savedStaticQubits); + }); + allocatedQubits.clear(); + allocatedQregs.clear(); + staticQubits.clear(); + + setInsertionPoint(mainFunc); + auto function = func::FuncOp::create( + *this, name, getFunctionType(argumentTypes, TypeRange{})); + function.setPrivate(); + auto* block = function.addEntryBlock(); + setInsertionPointToStart(block); + + SmallVector results = body(block->getArguments()); + if (block->mightHaveTerminator()) { + llvm::reportFatalUsageError( + "Function callback must not create a terminator"); + } + + for (Value result : results) { + allocatedQubits.remove(result); + allocatedQregs.remove(result); + } + for (Value qubit : allocatedQubits) { + DeallocOp::create(*this, qubit); + } + for (Value qreg : allocatedQregs) { + memref::DeallocOp::create(*this, qreg); + } + + function.setType( + getFunctionType(argumentTypes, ValueRange(results).getTypes())); + func::ReturnOp::create(*this, results); + return function; +} + +func::FuncOp QCProgramBuilder::createUnitaryFunction( + const StringRef name, const TypeRange argumentTypes, + const function_ref body) { + auto function = + createFunction(name, argumentTypes, [&](ValueRange arguments) { + body(arguments); + return SmallVector{}; + }); + mqt::setUnitaryFunction(function); + return function; +} + +SmallVector QCProgramBuilder::call(func::FuncOp callee, + ValueRange operands) { + checkFinalized(); + if (callee->getParentOp() != moduleOp_ || + callee.getArgumentTypes() != operands.getTypes()) { + llvm::reportFatalUsageError( + "Call operands must match a function in the current module"); + } + if (mqt::isUnitaryFunction(callee)) { + CallOp::create(*this, + FlatSymbolRefAttr::get(getContext(), callee.getName()), + operands); + return {}; + } + auto callOp = func::CallOp::create(*this, callee, operands); + return SmallVector(callOp.getResults()); +} + Value QCProgramBuilder::boolConstant(const bool value) { checkFinalized(); return arith::ConstantOp::create(*this, getBoolAttr(value)).getResult(); @@ -128,8 +217,15 @@ Value QCProgramBuilder::staticQubit(const uint64_t index) { } OpBuilder::InsertionGuard guard(*this); - auto mainFunc = mqt::getEntryPoint(cast(module)); - setInsertionPointToStart(&mainFunc.getBody().front()); + Operation* parent = getInsertionBlock()->getParentOp(); + auto function = dyn_cast(parent); + if (!function) { + function = parent->getParentOfType(); + } + if (!function) { + llvm::reportFatalInternalError("Static qubit has no enclosing function"); + } + setInsertionPointToStart(&function.getBody().front()); auto qubit = StaticOp::create(*this, index).getQubit(); staticQubits.try_emplace(index, qubit); return qubit; @@ -792,17 +888,12 @@ OwningOpRef QCProgramBuilder::finalize() { OwningOpRef QCProgramBuilder::finalize(ValueRange returnValues) { checkFinalized(); - // Ensure that main function exists and insertion point is valid + // Ensure that the entry-point function exists and the insertion point is + // valid. auto* insertionBlock = getInsertionBlock(); - func::FuncOp mainFunc = nullptr; - for (auto op : cast(module).getOps()) { - if (op.getName() == "main") { - mainFunc = op; - break; - } - } - if (!mainFunc) { - llvm::reportFatalUsageError("Could not find main function"); + auto mainFunc = mqt::getEntryPoint(cast(moduleOp_)); + if (mainFunc == nullptr) { + llvm::reportFatalUsageError("Could not find entry-point function"); } if ((insertionBlock == nullptr) || insertionBlock != &mainFunc.getBody().front()) { @@ -827,7 +918,7 @@ OwningOpRef QCProgramBuilder::finalize(ValueRange returnValues) { ctx = nullptr; // Transfer ownership to the caller - return cast(module); + return cast(moduleOp_); } OwningOpRef QCProgramBuilder::build( diff --git a/mlir/lib/Dialect/QC/IR/Operations/CallOp.cpp b/mlir/lib/Dialect/QC/IR/Operations/CallOp.cpp new file mode 100644 index 0000000000..d1dafeaa29 --- /dev/null +++ b/mlir/lib/Dialect/QC/IR/Operations/CallOp.cpp @@ -0,0 +1,64 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#include "mlir/Dialect/MQT/IR/MQTDialect.h" +#include "mlir/Dialect/QC/IR/QCDialect.h" +#include "mlir/Dialect/QC/IR/QCOps.h" + +#include +#include +#include +#include +#include + +#include +#include + +using namespace mlir; +using namespace mlir::qc; + +size_t CallOp::getNumParams() { + return static_cast(std::distance( + getOperands().begin(), llvm::find_if(getOperands(), [](Value value) { + return isa(value.getType()); + }))); +} + +size_t CallOp::getNumQubits() { return getNumOperands() - getNumParams(); } + +OperandRange CallOp::getParameters() { + return getOperands().take_front(getNumParams()); +} + +OperandRange CallOp::getQubits() { + return getOperands().drop_front(getNumParams()); +} + +LogicalResult CallOp::verifySymbolUses(SymbolTableCollection& symbolTable) { + auto function = + symbolTable.lookupNearestSymbolFrom(*this, getCalleeAttr()); + if (!function) { + return emitOpError() << "'" << getCallee() + << "' does not reference a valid function"; + } + if (!mqt::isUnitaryFunction(function)) { + return emitOpError() << "callee '" << getCallee() + << "' is not marked with mqt.unitary"; + } + if (function.getArgumentTypes() != getOperandTypes()) { + return emitOpError() << "operand types " << getOperandTypes() + << " do not match callee argument types " + << function.getArgumentTypes(); + } + if (function.getNumResults() != 0) { + return emitOpError("unitary QC callee must not return values"); + } + return success(); +} diff --git a/mlir/lib/Dialect/QCO/Builder/CMakeLists.txt b/mlir/lib/Dialect/QCO/Builder/CMakeLists.txt index d8f282e109..2cbaa1f729 100644 --- a/mlir/lib/Dialect/QCO/Builder/CMakeLists.txt +++ b/mlir/lib/Dialect/QCO/Builder/CMakeLists.txt @@ -19,7 +19,8 @@ add_mlir_library( MLIRQCODialect MLIRQTensorDialect PRIVATE - MLIRMQTUtils) + MLIRMQTUtils + MLIRQCOUtils) mqt_mlir_target_use_project_options(MLIRQCOProgramBuilder) diff --git a/mlir/lib/Dialect/QCO/Builder/QCOProgramBuilder.cpp b/mlir/lib/Dialect/QCO/Builder/QCOProgramBuilder.cpp index c0bab129df..cb5fa2de36 100644 --- a/mlir/lib/Dialect/QCO/Builder/QCOProgramBuilder.cpp +++ b/mlir/lib/Dialect/QCO/Builder/QCOProgramBuilder.cpp @@ -18,11 +18,13 @@ #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QCO/QCOUtils.h" +#include "mlir/Dialect/QCO/Utils/FunctionUtils.h" #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include #include +#include #include #include #include @@ -37,6 +39,7 @@ #include #include #include +#include #include #include #include @@ -55,7 +58,7 @@ namespace mlir::qco { QCOProgramBuilder::QCOProgramBuilder(MLIRContext* context) : ImplicitLocOpBuilder( FileLineColLoc::get(context, "", 1, 1), context), - ctx(context), module(ModuleOp::create(*this)) { + ctx(context), moduleOp_(ModuleOp::create(*this)) { ctx->loadDialect(); } @@ -64,7 +67,7 @@ void QCOProgramBuilder::initialize() { initialize({getI64Type()}); } void QCOProgramBuilder::initialize(TypeRange returnTypes) { // Set insertion point to the module body - setInsertionPointToStart(cast(module).getBody()); + setInsertionPointToStart(cast(moduleOp_).getBody()); // Create main function as entry point auto funcType = getFunctionType({}, returnTypes); @@ -78,7 +81,7 @@ void QCOProgramBuilder::initialize(TypeRange returnTypes) { } void QCOProgramBuilder::retype(TypeRange returnTypes) { - auto mainFunc = mqt::getEntryPoint(cast(module)); + auto mainFunc = mqt::getEntryPoint(cast(moduleOp_)); if (!mainFunc) { llvm::reportFatalUsageError("Main function not found for retyping"); } @@ -87,6 +90,167 @@ void QCOProgramBuilder::retype(TypeRange returnTypes) { mainFunc.setType(funcType); } +static bool isQubitTensor(Type type) { + auto tensor = dyn_cast(type); + return tensor && isa(tensor.getElementType()); +} + +func::FuncOp QCOProgramBuilder::createFunction( + const StringRef name, const TypeRange argumentTypes, + const function_ref(ValueRange)> body) { + checkFinalized(); + auto moduleOp = cast(moduleOp_); + auto mainFunc = mqt::getEntryPoint(moduleOp); + if (!mainFunc) { + llvm::reportFatalUsageError( + "QCOProgramBuilder must be initialized before creating a function"); + } + if (SymbolTable::lookupSymbolIn(moduleOp, name) != nullptr) { + llvm::reportFatalUsageError("Function name is already defined"); + } + + const InsertionGuard insertionGuard(*this); + auto savedQubits = std::move(validQubits); + auto savedTensors = std::move(validTensors); + const auto savedTensorCounter = tensorCounter; + auto stateGuard = llvm::make_scope_exit([&] { + validQubits = std::move(savedQubits); + validTensors = std::move(savedTensors); + tensorCounter = savedTensorCounter; + }); + validQubits.clear(); + validTensors.clear(); + + setInsertionPoint(mainFunc); + auto function = func::FuncOp::create( + *this, name, getFunctionType(argumentTypes, TypeRange{})); + function.setPrivate(); + auto* block = function.addEntryBlock(); + setInsertionPointToStart(block); + for (auto argument : block->getArguments()) { + if (isa(argument.getType())) { + validQubits.insert(argument); + } else if (isQubitTensor(argument.getType())) { + validTensors.insert(Tensor{argument, tensorCounter++}); + } + } + + SmallVector results = body(block->getArguments()); + if (block->mightHaveTerminator()) { + llvm::reportFatalUsageError( + "Function callback must not create a terminator"); + } + function.setType( + getFunctionType(argumentTypes, ValueRange(results).getTypes())); + SmallVector qubitArguments; + for (auto [index, argument] : llvm::enumerate(block->getArguments())) { + if (isa(argument.getType())) { + qubitArguments.emplace_back(index); + } + } + if (results.size() < qubitArguments.size()) { + llvm::reportFatalUsageError( + "Function must return every qubit argument as a trailing result"); + } + const auto firstQubitResult = results.size() - qubitArguments.size(); + if (llvm::any_of( + ValueRange(results).drop_front(firstQubitResult), + [](Value result) { return !isa(result.getType()); })) { + llvm::reportFatalUsageError( + "Function must return every qubit argument as a trailing result"); + } + for (auto [offset, argument] : llvm::enumerate(qubitArguments)) { + auto origin = + traceQubitArgument(function, results[firstQubitResult + offset]); + if (failed(origin) || *origin != argument) { + llvm::reportFatalUsageError( + "Function must return every qubit argument as a trailing result"); + } + } + for (auto [index, result] : llvm::enumerate(results)) { + if (isa(result.getType())) { + validateQubitValue(result); + validQubits.erase(result); + } else if (isQubitTensor(result.getType())) { + validateTensorValue(result); + validTensors.erase(result); + } + } + disposeLinearValues(); + func::ReturnOp::create(*this, results); + return function; +} + +func::FuncOp QCOProgramBuilder::createUnitaryFunction( + const StringRef name, const TypeRange argumentTypes, + const function_ref(ValueRange)> body) { + auto function = createFunction(name, argumentTypes, body); + mqt::setUnitaryFunction(function); + return function; +} + +SmallVector QCOProgramBuilder::call(func::FuncOp callee, + ValueRange operands) { + checkFinalized(); + if (callee->getParentOp() != moduleOp_ || + callee.getArgumentTypes() != operands.getTypes()) { + llvm::reportFatalUsageError( + "Call operands must match a function in the current module"); + } + if (llvm::any_of(operands, [](Value operand) { + return isQubitTensor(operand.getType()); + })) { + llvm::reportFatalUsageError( + "Quantum tensor function calls are not supported"); + } + + SmallVector qubitArguments; + for (auto operand : operands) { + if (!isa(operand.getType())) { + continue; + } + validateQubitValue(operand); + auto iterator = validQubits.find(operand); + qubitArguments.emplace_back(*iterator); + validQubits.erase(iterator); + } + + SmallVector results; + if (mqt::isUnitaryFunction(callee)) { + auto call = CallOp::create( + *this, FlatSymbolRefAttr::get(getContext(), callee.getName()), + operands); + llvm::append_range(results, call.getResults()); + } else { + auto call = func::CallOp::create(*this, callee, operands); + llvm::append_range(results, call.getResults()); + } + + if (results.size() < qubitArguments.size()) { + llvm::reportFatalUsageError( + "Callee does not return its qubit arguments positionally"); + } + const auto firstQubitResult = results.size() - qubitArguments.size(); + if (llvm::any_of( + ValueRange(results).drop_front(firstQubitResult), + [](Value result) { return !isa(result.getType()); })) { + llvm::reportFatalUsageError( + "Callee does not return its qubit arguments positionally"); + } + for (auto [index, result] : llvm::enumerate(results)) { + if (!isa(result.getType())) { + continue; + } + if (index >= firstQubitResult) { + const auto& tracked = qubitArguments[index - firstQubitResult]; + validQubits.insert(Qubit{result, tracked.regId, tracked.regIndex}); + } else { + validQubits.insert(result); + } + } + return results; +} + Value QCOProgramBuilder::intConstant(const int64_t value) { checkFinalized(); return arith::ConstantOp::create(*this, getI64IntegerAttr(value)).getResult(); @@ -1443,6 +1607,33 @@ void QCOProgramBuilder::ensureAllocationMode( llvm::reportFatalUsageError(message.c_str()); } +void QCOProgramBuilder::disposeLinearValues() { + DenseSet validTensorIds; + for (const auto& tensor : validTensors) { + validTensorIds.insert(tensor.regId); + } + + DenseMap> qubitsByRegister; + for (const auto& qubit : validQubits) { + if (qubit.regId == -1 || !validTensorIds.contains(qubit.regId)) { + SinkOp::create(*this, qubit); + } else { + qubitsByRegister[qubit.regId].emplace_back(qubit); + } + } + for (const auto& tensor : validTensors) { + Value currentTensor = tensor; + for (const auto& qubit : qubitsByRegister[tensor.regId]) { + currentTensor = + qtensor::InsertOp::create(*this, qubit, currentTensor, qubit.regIndex) + .getResult(); + } + qtensor::DeallocOp::create(*this, currentTensor); + } + validQubits.clear(); + validTensors.clear(); +} + OwningOpRef QCOProgramBuilder::finalize() { checkFinalized(); @@ -1453,17 +1644,12 @@ OwningOpRef QCOProgramBuilder::finalize() { OwningOpRef QCOProgramBuilder::finalize(ValueRange returnValues) { checkFinalized(); - // Ensure that main function exists and insertion point is valid + // Ensure that the entry-point function exists and the insertion point is + // valid. auto* insertionBlock = getInsertionBlock(); - func::FuncOp mainFunc = nullptr; - for (auto op : cast(module).getOps()) { - if (op.getName() == "main") { - mainFunc = op; - break; - } - } - if (!mainFunc) { - llvm::reportFatalUsageError("Could not find main function"); + auto mainFunc = mqt::getEntryPoint(cast(moduleOp_)); + if (mainFunc == nullptr) { + llvm::reportFatalUsageError("Could not find entry-point function"); } if ((insertionBlock == nullptr) || insertionBlock != &mainFunc.getBody().front()) { @@ -1484,35 +1670,7 @@ OwningOpRef QCOProgramBuilder::finalize(ValueRange returnValues) { } } - DenseSet validTensorIds; - for (const auto& tensor : validTensors) { - validTensorIds.insert(tensor.regId); - } - - DenseMap> qubitsByRegister; - for (const auto& qubit : validQubits) { - if (qubit.regId == -1 || !validTensorIds.contains(qubit.regId)) { - // Automatically deallocate all still-allocated qubits - SinkOp::create(*this, qubit); - } else { - qubitsByRegister[qubit.regId].emplace_back(qubit); - } - } - - // Automatically deallocate all still-allocated tensors - for (const auto& tensor : validTensors) { - Value currentTensor = tensor; - // Filter out qubits belonging to this tensor - for (const auto& qubit : qubitsByRegister[tensor.regId]) { - currentTensor = - qtensor::InsertOp::create(*this, qubit, currentTensor, qubit.regIndex) - .getResult(); - } - // Deallocate tensor - qtensor::DeallocOp::create(*this, currentTensor); - } - validQubits.clear(); - validTensors.clear(); + disposeLinearValues(); // Add return statement with the given return values to the main function func::ReturnOp::create(*this, returnValues); @@ -1520,7 +1678,7 @@ OwningOpRef QCOProgramBuilder::finalize(ValueRange returnValues) { // Invalidate context to prevent use-after-finalize ctx = nullptr; - return cast(module); + return cast(moduleOp_); } OwningOpRef QCOProgramBuilder::build( diff --git a/mlir/lib/Dialect/QCO/IR/Operations/CallOp.cpp b/mlir/lib/Dialect/QCO/IR/Operations/CallOp.cpp new file mode 100644 index 0000000000..1338ac3482 --- /dev/null +++ b/mlir/lib/Dialect/QCO/IR/Operations/CallOp.cpp @@ -0,0 +1,106 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#include "mlir/Dialect/MQT/IR/MQTDialect.h" +#include "mlir/Dialect/QCO/IR/QCOOps.h" + +#include +#include +#include +#include +#include + +#include +#include + +using namespace mlir; +using namespace mlir::qco; + +void CallOp::build(OpBuilder&, OperationState& state, FlatSymbolRefAttr callee, + ValueRange operands) { + state.addAttribute("callee", callee); + state.addOperands(operands); + for (Value operand : operands) { + if (isa(operand.getType())) { + state.addTypes(operand.getType()); + } + } +} + +size_t CallOp::getNumParams() { + return getNumOperands() < getNumResults() + ? 0 + : getNumOperands() - getNumResults(); +} + +OperandRange CallOp::getParameters() { + return getOperands().take_front(getNumParams()); +} + +OperandRange CallOp::getInputQubits() { + return getOperands().drop_front(getNumParams()); +} + +Value CallOp::getInputForOutput(Value output) { + auto result = dyn_cast(output); + auto inputs = getInputQubits(); + if (!result || result.getOwner() != getOperation() || + result.getResultNumber() >= inputs.size()) { + return {}; + } + return inputs[result.getResultNumber()]; +} + +Value CallOp::getOutputForInput(Value input) { + const auto position = llvm::find(getInputQubits(), input); + if (position == getInputQubits().end()) { + return {}; + } + return getOutputQubit( + static_cast(std::distance(getInputQubits().begin(), position))); +} + +LogicalResult CallOp::verify() { + if (getNumOperands() < getNumResults() || + llvm::any_of( + getParameters(), + [](Value value) { return isa(value.getType()); }) || + llvm::any_of(getInputQubits(), [](Value value) { + return !isa(value.getType()); + })) { + return emitOpError( + "requires one trailing qubit operand for every qubit result"); + } + return success(); +} + +LogicalResult CallOp::verifySymbolUses(SymbolTableCollection& symbolTable) { + auto function = + symbolTable.lookupNearestSymbolFrom(*this, getCalleeAttr()); + if (!function) { + return emitOpError() << "'" << getCallee() + << "' does not reference a valid function"; + } + if (!mqt::isUnitaryFunction(function)) { + return emitOpError() << "callee '" << getCallee() + << "' is not marked with mqt.unitary"; + } + if (function.getArgumentTypes() != getOperandTypes()) { + return emitOpError() << "operand types " << getOperandTypes() + << " do not match callee argument types " + << function.getArgumentTypes(); + } + if (function.getResultTypes() != getResultTypes()) { + return emitOpError() << "result types " << getResultTypes() + << " do not match callee result types " + << function.getResultTypes(); + } + return success(); +} diff --git a/mlir/lib/Dialect/QCO/Utils/FunctionUtils.cpp b/mlir/lib/Dialect/QCO/Utils/FunctionUtils.cpp new file mode 100644 index 0000000000..3523e16e8e --- /dev/null +++ b/mlir/lib/Dialect/QCO/Utils/FunctionUtils.cpp @@ -0,0 +1,66 @@ +/* + * Copyright (c) 2023 - 2026 Chair for Design Automation, TUM + * Copyright (c) 2025 - 2026 Munich Quantum Software Company GmbH + * All rights reserved. + * + * SPDX-License-Identifier: MIT + * + * Licensed under the MIT License + */ + +#include "mlir/Dialect/QCO/Utils/FunctionUtils.h" + +#include "mlir/Dialect/QCO/IR/QCODialect.h" +#include "mlir/Dialect/QCO/Utils/WireIterator.h" + +#include +#include +#include + +using namespace mlir; +using namespace mlir::qco; + +FailureOr mlir::qco::traceQubitArgument(func::FuncOp function, + Value value) { + if (function.isDeclaration()) { + return failure(); + } + while (true) { + if (auto argument = dyn_cast(value)) { + if (argument.getOwner() == &function.getBody().front() && + isa(argument.getType())) { + return argument.getArgNumber(); + } + return failure(); + } + + if (auto call = value.getDefiningOp()) { + auto callee = SymbolTable::lookupNearestSymbolFrom( + call, call.getCalleeAttr()); + if (!callee) { + return failure(); + } + SmallVector qubitArguments; + for (auto [index, type] : llvm::enumerate(callee.getArgumentTypes())) { + if (isa(type)) { + qubitArguments.emplace_back(index); + } + } + auto result = cast(value).getResultNumber(); + if (call.getNumResults() < qubitArguments.size() || + result < call.getNumResults() - qubitArguments.size()) { + return failure(); + } + value = call.getOperand(qubitArguments[result - (call.getNumResults() - + qubitArguments.size())]); + continue; + } + + WireIterator iterator(value); + --iterator; + if (iterator == std::default_sentinel) { + return failure(); + } + value = iterator.qubit(); + } +} diff --git a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp index 774545b5c9..af090c2ef6 100644 --- a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp +++ b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp @@ -16,176 +16,21 @@ #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include -#include #include #include #include #include #include -#include #include #include #include #include #include -#include #include -#include -#include namespace mlir::qco { -// Returns the position of a qubit among the qubit-typed values in a range. -template -static std::optional qubitPositionIn(RangeT range, Value qubit) { - size_t position = 0; - for (Value value : range) { - if (!isa(value.getType())) { - continue; - } - if (value == qubit) { - return position; - } - ++position; - } - return std::nullopt; -} - -// Returns the qubit-typed value at a position, or null if none exists. -template -static Value nthQubitOf(RangeT range, size_t position) { - size_t seen = 0; - for (Value value : range) { - if (!isa(value.getType())) { - continue; - } - if (seen == position) { - return value; - } - ++seen; - } - return nullptr; -} - -FailureOr> -CallQubitMapping::computeMapping(func::FuncOp callee) { - if (callee.isExternal()) { - return failure(); - } - - // Threading a callee already in progress would not terminate. - if (!inProgress.insert(callee.getOperation()).second) { - return failure(); - } - auto progressGuard = - llvm::make_scope_exit([&] { inProgress.erase(callee.getOperation()); }); - - // A body under construction may not have a terminator yet. - if (!callee.getBody().hasOneBlock() || - !callee.getBody().front().mightHaveTerminator()) { - return failure(); - } - auto returnOp = - dyn_cast(callee.getBody().front().getTerminator()); - if (!returnOp) { - return failure(); - } - - SmallVector mapping; - for (BlockArgument arg : callee.getArguments()) { - if (!isa(arg.getType())) { - continue; - } - - int64_t resultIndex = KEPT; - { - // Follow the argument to the end of its wire. - Value last = arg; - Operation* lastOp = nullptr; - WireIterator it(arg, this); - for (; it != std::default_sentinel; ++it) { - last = it.qubit(); - lastOp = it.operation(); - } - if (it.mappingFailed_) { - return failure(); - } - - if (isa_and_nonnull(lastOp)) { - for (const auto& [index, operand] : - llvm::enumerate(returnOp.getOperands())) { - if (operand == last) { - resultIndex = static_cast(index); - break; - } - } - } - } - mapping.emplace_back(resultIndex); - } - - return mapping; -} - -void CallQubitMapping::invalidate() { cache.clear(); } - -FailureOr> CallQubitMapping::mappingFor(func::CallOp callOp) { - auto callee = dyn_cast_or_null( - SymbolTable::lookupNearestSymbolFrom(callOp, callOp.getCalleeAttr())); - if (!callee) { - return failure(); - } - - auto* const key = callee.getOperation(); - if (const auto it = cache.find(key); it != cache.end()) { - return ArrayRef(it->second); - } - // Compute before caching so recursion is detected through inProgress. - auto mapping = computeMapping(callee); - if (failed(mapping)) { - return failure(); - } - return ArrayRef( - cache.insert_or_assign(key, std::move(*mapping)).first->second); -} - -FailureOr CallQubitMapping::getResultForOperand(func::CallOp callOp, - Value operand) { - const auto position = qubitPositionIn(callOp.getOperands(), operand); - assert(position && "expected a qubit operand of the call"); - auto mappingOr = mappingFor(callOp); - if (failed(mappingOr)) { - return failure(); - } - ArrayRef mapping = *mappingOr; - assert(*position < mapping.size() && "expected matching call signature"); - const auto resultIndex = mapping[*position]; - if (resultIndex == KEPT) { - return Value{}; - } - return callOp.getResult(static_cast(resultIndex)); -} - -FailureOr CallQubitMapping::getOperandForResult(func::CallOp callOp, - Value result) { - auto opResult = cast(result); - assert(opResult.getOwner() == callOp.getOperation() && - "expected a result of the call"); - const auto resultIndex = static_cast(opResult.getResultNumber()); - auto mappingOr = mappingFor(callOp); - if (failed(mappingOr)) { - return failure(); - } - ArrayRef mapping = *mappingOr; - for (const auto& [position, index] : llvm::enumerate(mapping)) { - if (index == resultIndex) { - return nthQubitOf(callOp.getOperands(), position); - } - } - return Value{}; -} - bool WireIterator::isTail(Operation* op) { // `qtensor.from_elements` takes qubits into a tensor just like // `qtensor.insert` does, so a wire reaching either of them ends there. @@ -205,20 +50,6 @@ Operation* WireIterator::operation() const { return op_; } -FailureOr WireIterator::resultForOperand(func::CallOp callOp, - Value operand) const { - CallQubitMapping local; - auto& mapping = mapping_ == nullptr ? local : *mapping_; - return mapping.getResultForOperand(callOp, operand); -} - -Value WireIterator::operandForResult(func::CallOp callOp, Value result) const { - CallQubitMapping local; - auto& mapping = mapping_ == nullptr ? local : *mapping_; - auto operand = mapping.getOperandForResult(callOp, result); - return succeeded(operand) ? *operand : Value{}; -} - Value WireIterator::qubit() const { if (*this == std::default_sentinel) { llvm::reportFatalInternalError("Trying to access qubit of sentinel!"); @@ -278,26 +109,7 @@ void WireIterator::forward() { .Case([&](IndexSwitchOp op) { qubit_ = op.getTiedResult(&(*qubit_.use_begin())); }) - .Case([&](func::CallOp op) { - // A call threads the qubit through to the matching result. When the - // callee keeps it, the wire ends here. - - auto result = resultForOperand(op, qubit_); - if (failed(result)) { - mappingFailed_ = true; - pos_ = Position::Tail; - return; - } - if (!*result) { - pos_ = Position::Tail; - return; - } - qubit_ = *result; - }) - .Default([&](Operation*) { - mappingFailed_ = true; - pos_ = Position::Tail; - }); + .Default([&](Operation*) { pos_ = Position::Tail; }); } void WireIterator::backward() { @@ -366,14 +178,6 @@ void WireIterator::backward() { } llvm::reportFatalInternalError("expected result lookup"); }) - .Case([&](func::CallOp callOp) { - Value operand = operandForResult(callOp, qubit_); - if (!operand) { - unknown = true; - return; - } - qubit_ = operand; - }) .Default([&](Operation*) { unknown = true; }); if (unknown) { diff --git a/mlir/lib/Dialect/QTensor/Utils/TensorIterator.cpp b/mlir/lib/Dialect/QTensor/Utils/TensorIterator.cpp index c255e22fc7..0a6914841e 100644 --- a/mlir/lib/Dialect/QTensor/Utils/TensorIterator.cpp +++ b/mlir/lib/Dialect/QTensor/Utils/TensorIterator.cpp @@ -15,23 +15,18 @@ #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include -#include #include #include #include #include #include -#include #include #include #include #include #include -#include #include -#include -#include namespace mlir::qtensor { TypedValue TensorIterator::tensor() const { @@ -235,148 +230,4 @@ void TensorIterator::backward() { static_assert(std::bidirectional_iterator); static_assert(std::sentinel_for, "std::default_sentinel_t must be a sentinel for TensorIterator."); - -// Returns whether a type is a tensor of qubits. -static bool isQubitTensor(Type type) { - auto tensorType = dyn_cast(type); - return tensorType && isa(tensorType.getElementType()); -} - -// Returns the position of a value among the qubit tensors in a range. -static std::optional tensorPositionIn(ValueRange range, Value value) { - size_t position = 0; - for (Value candidate : range) { - if (!isQubitTensor(candidate.getType())) { - continue; - } - if (candidate == value) { - return position; - } - ++position; - } - return std::nullopt; -} - -FailureOr CallTensorMapping::threadToResult(Value arg, - func::ReturnOp returnOp) { - Value current = arg; - while (true) { - // Follow the chain to its end. `tensor()` is null on the operations that - // consume a tensor without producing one, so the last non-null value is - // the one the terminating operation takes. - Value last = current; - Operation* lastOp = nullptr; - for (TensorIterator it(cast>(current)); - it != std::default_sentinel; ++it) { - if (Value currentTensor = it.tensor()) { - last = currentTensor; - } - lastOp = it.operation(); - } - - if (isa_and_nonnull(lastOp)) { - for (const auto& [index, operand] : - llvm::enumerate(returnOp.getOperands())) { - if (operand == last) { - return static_cast(index); - } - } - return KEPT; - } - - // The chain stops at a nested call. Step over it to the result that - // continues the tensor and keep following from there. Each hop moves - // forward along the def-use chain, so this terminates. - auto callOp = dyn_cast_or_null(lastOp); - if (!callOp) { - return KEPT; - } - auto next = getResultForOperand(callOp, last); - if (failed(next)) { - return failure(); - } - if (!*next) { - return KEPT; - } - current = *next; - } -} - -FailureOr> -CallTensorMapping::computeMapping(func::FuncOp callee) { - if (callee.isExternal()) { - return failure(); - } - - // Threading a callee already in progress would not terminate. - if (!inProgress.insert(callee.getOperation()).second) { - return failure(); - } - auto progressGuard = - llvm::make_scope_exit([&] { inProgress.erase(callee.getOperation()); }); - - // A body under construction may not have a terminator yet. - if (!callee.getBody().hasOneBlock() || - !callee.getBody().front().mightHaveTerminator()) { - return failure(); - } - auto returnOp = - dyn_cast(callee.getBody().front().getTerminator()); - if (!returnOp) { - return failure(); - } - - SmallVector mapping; - for (BlockArgument arg : callee.getArguments()) { - if (!isQubitTensor(arg.getType())) { - continue; - } - auto result = threadToResult(arg, returnOp); - if (failed(result)) { - return failure(); - } - mapping.emplace_back(*result); - } - - return mapping; -} - -FailureOr> -CallTensorMapping::mappingFor(func::CallOp callOp) { - auto callee = dyn_cast_or_null( - SymbolTable::lookupNearestSymbolFrom(callOp, callOp.getCalleeAttr())); - if (!callee) { - return failure(); - } - - auto* const key = callee.getOperation(); - if (const auto it = cache.find(key); it != cache.end()) { - return ArrayRef(it->second); - } - // Compute before caching so recursion is detected through inProgress. - auto mapping = computeMapping(callee); - if (failed(mapping)) { - return failure(); - } - return ArrayRef( - cache.insert_or_assign(key, std::move(*mapping)).first->second); -} - -FailureOr CallTensorMapping::getResultForOperand(func::CallOp callOp, - Value operand) { - const auto position = tensorPositionIn(callOp.getOperands(), operand); - assert(position && "expected a qubit-tensor operand of the call"); - auto mappingOr = mappingFor(callOp); - if (failed(mappingOr)) { - return failure(); - } - ArrayRef mapping = *mappingOr; - assert(*position < mapping.size() && "expected matching call signature"); - const auto resultIndex = mapping[*position]; - if (resultIndex == KEPT) { - return Value{}; - } - return callOp.getResult(static_cast(resultIndex)); -} - } // namespace mlir::qtensor diff --git a/mlir/unittests/Compiler/test_compiler_pipeline.cpp b/mlir/unittests/Compiler/test_compiler_pipeline.cpp index 8e5a3724af..e1d333f4f6 100644 --- a/mlir/unittests/Compiler/test_compiler_pipeline.cpp +++ b/mlir/unittests/Compiler/test_compiler_pipeline.cpp @@ -134,17 +134,17 @@ class CompilerPipelineTest [[nodiscard]] OwningOpRef buildQCReference(const QCProgramBuilderFn builder) const { - auto module = ::mqt::test::buildMLIRProgram(context.get(), builder); - EXPECT_TRUE(runQCCleanupPipeline(module.get()).succeeded()); - return module; + auto moduleOp = ::mqt::test::buildMLIRProgram(context.get(), builder); + EXPECT_TRUE(runQCCleanupPipeline(moduleOp.get()).succeeded()); + return moduleOp; } [[nodiscard]] OwningOpRef buildQIRReference(const QIRProgramBuilderFn builder) const { - auto module = ::mqt::test::buildMLIRProgram( + auto moduleOp = ::mqt::test::buildMLIRProgram( context.get(), builder, QIRProgramBuilder::Profile::Adaptive); - EXPECT_TRUE(runQIRCleanupPipeline(module.get(), true).succeeded()); - return module; + EXPECT_TRUE(runQIRCleanupPipeline(moduleOp.get(), true).succeeded()); + return moduleOp; } [[nodiscard]] OwningOpRef @@ -152,16 +152,16 @@ class CompilerPipelineTest return parseSourceString(ir, context.get()); } - static void ignoreSingleQIRResultLabel(ModuleOp module) { + static void ignoreSingleQIRResultLabel(ModuleOp moduleOp) { constexpr llvm::StringLiteral prefix = "qir.result_label_"; size_t numLabels = 0; - module.walk([&](LLVM::GlobalOp op) { + moduleOp.walk([&](LLVM::GlobalOp op) { numLabels += op.getSymName().starts_with(prefix); }); if (numLabels != 1) { return; } - module.walk([&](Operation* op) { + moduleOp.walk([&](Operation* op) { if (const auto name = op->getAttrOfType("sym_name"); name && name.getValue().starts_with(prefix)) { op->removeAttr("sym_name"); @@ -233,15 +233,15 @@ TEST_P(CompilerPipelineTest, EndToEndPipeline) { DeferredPrinter printer; ASSERT_TRUE(testCase.qcProgramBuilder); - auto module = + auto moduleOp = ::mqt::test::buildMLIRProgram(context.get(), testCase.qcProgramBuilder); - ASSERT_TRUE(module); - printer.record(module.get(), "QC Input" + name); - EXPECT_TRUE(verify(*module).succeeded()); + ASSERT_TRUE(moduleOp); + printer.record(moduleOp.get(), "QC Input" + name); + EXPECT_TRUE(verify(*moduleOp).succeeded()); std::string source; llvm::raw_string_ostream sourceStream(source); - module->print(sourceStream); + moduleOp->print(sourceStream); auto input = QCProgram::fromMLIRString(source); ASSERT_TRUE(input); auto compiled = runDefaultPipeline( @@ -419,7 +419,7 @@ TEST(CompilerProgramOwnershipTest, EnforcesQCOLinearityAtPublicBoundaries) { ProgramFormat::QCO)); } -/// Raw QCO stops before the registered default optimization pipeline. +// Raw QCO stops before the registered default optimization pipeline. TEST_F(CompilerPipelineTest, RawAndOptimizedQCOAreDistinctCheckpoints) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -458,7 +458,7 @@ h q; EXPECT_FALSE(std::get(*result).str().empty()); } -/// Test: typed programs transfer ownership between compiler dialects +// Test: typed programs transfer ownership between compiler dialects TEST_F(CompilerPipelineTest, TypedProgramsComposeWithoutImplicitCopies) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -995,7 +995,7 @@ INSTANTIATE_TEST_SUITE_P(OpenQASMPrograms, OpenQASMJeffBoundaryTest, } // namespace -/// Test: typed programs import MLIR and OpenQASM from their public APIs +// Test: typed programs import MLIR and OpenQASM from their public APIs TEST_F(CompilerPipelineTest, TypedProgramImportsAndCopies) { const std::string mlir = R"(module { %0 = qc.alloc : !qc.qubit @@ -1037,7 +1037,7 @@ h q; EXPECT_FALSE(QCOProgram::fromMLIRString(mlir)); } -/// Test: QCO imports require each linear value to have one use. +// Test: QCO imports require each linear value to have one use. TEST_F(CompilerPipelineTest, QCOProgramImportsEnforceLinearity) { const std::string valid = R"mlir(module { func.func @main() { @@ -1084,7 +1084,7 @@ TEST_F(CompilerPipelineTest, QCOProgramImportsEnforceLinearity) { EXPECT_FALSE(QCOProgram::fromMLIRFile(path)); } -/// Test: typed programs emit OpenQASM directly and through the pipeline. +// Test: typed programs emit OpenQASM directly and through the pipeline. TEST_F(CompilerPipelineTest, TypedProgramsEmitOpenQASM) { const std::string qasm = R"(OPENQASM 3.1; include "stdgates.inc"; @@ -1153,7 +1153,7 @@ TEST_F(CompilerPipelineTest, TypedOpenQASMExportReportsUnsupportedQC) { EXPECT_FALSE(program->toOpenQASM3()); } -/// Test: typed programs expose idempotent global-phase normalization. +// Test: typed programs expose idempotent global-phase normalization. TEST_F(CompilerPipelineTest, TypedProgramsNormalizeGlobalPhases) { const std::string qcSource = R"mlir(module { func.func @test(%q: !qc.qubit) { @@ -1195,7 +1195,7 @@ TEST_F(CompilerPipelineTest, TypedProgramsNormalizeGlobalPhases) { EXPECT_EQ(StringRef(textual->str()).count("qco.gphase"), 1); } -/// Test: jeff programs round-trip through their binary APIs +// Test: jeff programs round-trip through their binary APIs TEST_F(CompilerPipelineTest, JeffProgramsRoundTripThroughBytesAndFiles) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1236,7 +1236,7 @@ x q; EXPECT_FALSE(jeff.write(path.parent_path() / "missing" / "output.jeff")); } -/// Test: QCO and QIR typed programs retain their respective semantics +// Test: QCO and QIR typed programs retain their respective semantics TEST_F(CompilerPipelineTest, QCOAndQIRProgramsImportCopyAndOptimize) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1290,7 +1290,7 @@ h q; base->writeBitcode(bitcodePath.parent_path() / "missing" / "output.bc")); } -/// Test: QCO program APIs configure and execute their associated passes. +// Test: QCO program APIs configure and execute their associated passes. TEST_F(CompilerPipelineTest, QCOProgramOptimizationAPIs) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1323,7 +1323,7 @@ cx q[0], q[2]; EXPECT_EQ(loopProgram->str().find("scf.for"), std::string::npos); } -/// Test: target compilation decomposes, maps, synthesizes, and verifies. +// Test: target compilation decomposes, maps, synthesizes, and verifies. TEST_F(CompilerPipelineTest, QCOProgramCompilesForTarget) { auto qc = QCProgram::fromQASMString(qasm::multipleControlledX); ASSERT_TRUE(qc); @@ -1367,7 +1367,7 @@ TEST_F(CompilerPipelineTest, QCOProgramCompilesForTarget) { EXPECT_FALSE(unsupportedQCO->compileForTarget(makeSparseUCZTarget(false))); } -/// Test that target compilation leaves dead-value cleanup at a fixed point. +// Test that target compilation leaves dead-value cleanup at a fixed point. TEST_F(CompilerPipelineTest, TargetCompilationLeavesDeadValueCleanupAtFixedPoint) { constexpr llvm::StringLiteral source = R"mlir( @@ -1603,7 +1603,7 @@ TEST_F(CompilerPipelineTest, QCOProgramMergesDynamicRunInNativeCtrlBody) { EXPECT_FALSE(main.getArgument(0).use_empty()); } -/// Test: all-to-all target compilation uses compact placement. +// Test: all-to-all target compilation uses compact placement. TEST_F(CompilerPipelineTest, QCOProgramUsesCompactAllToAllPlacement) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1645,7 +1645,7 @@ c = measure q; EXPECT_EQ(numSwaps, 0); } -/// Test: target compilation retains unobserved quantum operations. +// Test: target compilation retains unobserved quantum operations. TEST_F(CompilerPipelineTest, QCOProgramPreservesUnobservedQuantumOperations) { constexpr llvm::StringLiteral source = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1663,14 +1663,14 @@ h q[1]; auto program = std::move(*qc).intoQCO(); ASSERT_TRUE(program); ASSERT_TRUE(program->compileForTarget(target)); - auto module = parseRecordedModule(program->str()); - ASSERT_TRUE(module); - EXPECT_TRUE(verify(*module).succeeded()); + auto moduleOp = parseRecordedModule(program->str()); + ASSERT_TRUE(moduleOp); + EXPECT_TRUE(verify(*moduleOp).succeeded()); size_t unitaryOperations = 0; size_t resets = 0; size_t staticQubits = 0; - module->walk([&](Operation* operation) { + moduleOp->walk([&](Operation* operation) { unitaryOperations += isa(operation); resets += isa(operation); staticQubits += isa(operation); @@ -1680,7 +1680,7 @@ h q[1]; EXPECT_EQ(staticQubits, 2U); } -/// Test: the default pipeline accepts an optional compiler target. +// Test: the default pipeline accepts an optional compiler target. TEST_F(CompilerPipelineTest, DefaultPipelineCompilesForTarget) { auto input = QCProgram::fromQASMString(qasm::multipleControlledX); ASSERT_TRUE(input); @@ -1719,17 +1719,17 @@ TEST_F(CompilerPipelineTest, DefaultPipelineCompilesForTarget) { EXPECT_TRUE(qir.llvmIR()); } -/// Test: QCO programs expose the raw and composite qubit-reuse flows. +// Test: QCO programs expose the raw and composite qubit-reuse flows. TEST_F(CompilerPipelineTest, QCOProgramQubitReuseAPIs) { const auto countAllocations = [](const QCOProgram& program) { const auto ir = program.str(); return StringRef(ir).count("qco.alloc"); }; const auto buildQCO = [this](const QCProgramBuilderFn& builder) { - auto module = ::mqt::test::buildMLIRProgram(context.get(), builder); + auto moduleOp = ::mqt::test::buildMLIRProgram(context.get(), builder); std::string source; llvm::raw_string_ostream stream(source); - module->print(stream); + moduleOp->print(stream); auto qc = QCProgram::fromMLIRString(source); if (!qc) { return std::optional{}; @@ -1753,7 +1753,7 @@ TEST_F(CompilerPipelineTest, QCOProgramQubitReuseAPIs) { EXPECT_NE(compositeQCO->str().find("qco.reset"), std::string::npos); } -/// Test: default compilation returns the requested typed program format +// Test: default compilation returns the requested typed program format TEST_F(CompilerPipelineTest, DefaultPipelineSelectsRequestedProgramFormats) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -1858,17 +1858,17 @@ h q; EXPECT_TRUE(std::holds_alternative(*fromJeff)); } -/// Test: QCOProgram::decomposeMultiControlled runs the pass on MCX. -/// -/// Correctness of the decomposition is tested in a dedicated suite. +// Test: QCOProgram::decomposeMultiControlled runs the pass on MCX. +// +// Correctness of the decomposition is tested in a dedicated suite. TEST_F(CompilerPipelineTest, DecomposeMultiControlledPass) { - auto module = mlir::qc::QCProgramBuilder::build( + auto moduleOp = mlir::qc::QCProgramBuilder::build( context.get(), mlir::qc::multipleControlledX); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); std::string source; llvm::raw_string_ostream stream(source); - module->print(stream); + moduleOp->print(stream); auto input = QCProgram::fromMLIRString(source); ASSERT_TRUE(input); auto qco = std::move(*input).intoQCO(); @@ -1880,13 +1880,13 @@ TEST_F(CompilerPipelineTest, DecomposeMultiControlledPass) { } TEST_F(CompilerPipelineTest, DecomposeMultiControlledPassMcz) { - auto module = mlir::qc::QCProgramBuilder::build( + auto moduleOp = mlir::qc::QCProgramBuilder::build( context.get(), mlir::qc::multipleControlledZ); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); std::string source; llvm::raw_string_ostream stream(source); - module->print(stream); + moduleOp->print(stream); auto input = QCProgram::fromMLIRString(source); ASSERT_TRUE(input); auto qco = std::move(*input).intoQCO(); @@ -1902,12 +1902,12 @@ TEST_F(CompilerPipelineTest, EXPECT_FALSE(isDecomposeMultiControlledConfigValid(2U)); EXPECT_TRUE(isDecomposeMultiControlledConfigValid(3U)); - auto module = mlir::qc::QCProgramBuilder::build( + auto moduleOp = mlir::qc::QCProgramBuilder::build( context.get(), mlir::qc::multipleControlledX); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); std::string source; llvm::raw_string_ostream stream(source); - module->print(stream); + moduleOp->print(stream); auto input = QCProgram::fromMLIRString(source); ASSERT_TRUE(input); auto qco = std::move(*input).intoQCO(); @@ -1916,26 +1916,26 @@ TEST_F(CompilerPipelineTest, } TEST_F(CompilerPipelineTest, PopulateDecomposeMultiControlledPipeline) { - auto module = + auto moduleOp = QCOProgramBuilder::build(context.get(), [](QCOProgramBuilder& builder) { builder.mcx({builder.staticQubit(0), builder.staticQubit(1), builder.staticQubit(2)}, builder.staticQubit(3)); return SmallVector{}; }); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); std::string before; llvm::raw_string_ostream beforeStream(before); - module->print(beforeStream); + moduleOp->print(beforeStream); - PassManager pm(module->getContext()); + PassManager pm(moduleOp->getContext()); populateDecomposeMultiControlledPipeline(pm, 3); - ASSERT_TRUE(pm.run(module.get()).succeeded()); + ASSERT_TRUE(pm.run(moduleOp.get()).succeeded()); std::string after; llvm::raw_string_ostream afterStream(after); - module->print(afterStream); + moduleOp->print(afterStream); EXPECT_NE(after, before); } @@ -1990,6 +1990,15 @@ INSTANTIATE_TEST_SUITE_P( "HWithoutRegister", MQT_NAMED_BUILDER(mlir::qc::hWithoutRegister), MQT_NAMED_BUILDER(mlir::qc::hWithoutRegister), MQT_NAMED_BUILDER(mlir::qir::hWithoutRegister)}, + CompilerPipelineTestCase{ + "ReusableUnitaryFunction", + MQT_NAMED_BUILDER(mlir::qc::reusableUnitaryFunction), + MQT_NAMED_BUILDER(mlir::qc::reusableUnitaryFunction), nullptr, + false}, + CompilerPipelineTestCase{ + "ReusableResetFunction", + MQT_NAMED_BUILDER(mlir::qc::reusableResetFunction), + MQT_NAMED_BUILDER(mlir::qc::reusableResetFunction), nullptr, false}, CompilerPipelineTestCase{ "InverseIswap", MQT_NAMED_BUILDER(mlir::qc::inverseIswap), MQT_NAMED_BUILDER(mlir::qc::inverseIswap), nullptr, false}, @@ -2009,7 +2018,7 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(mlir::qir::singleControlledXOnIndividualQubits), true, "reuse-qubits,mqt-qco-default"})); -/// Test: gate counting respects modifiers and skips barriers. +// Test: gate counting respects modifiers and skips barriers. TEST_F(CompilerPipelineTest, QCProgramCountGates) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; @@ -2045,8 +2054,8 @@ TEST_F(CompilerPipelineTest, QCProgramCountGatesWithoutEntryPoint) { EXPECT_EQ(qc->numTwoQubitGates(), 0); } -/// Test: gate counting includes each structured control-flow region -/// once. +// Test: gate counting includes each structured control-flow region +// once. TEST_F(CompilerPipelineTest, QCProgramCountGatesInStructuredControlFlow) { const std::string qasm = R"(OPENQASM 3.0; include "stdgates.inc"; diff --git a/mlir/unittests/Conversion/QCOToQC/test_qco_to_qc.cpp b/mlir/unittests/Conversion/QCOToQC/test_qco_to_qc.cpp index c1c720b453..9c5bd2b5ab 100644 --- a/mlir/unittests/Conversion/QCOToQC/test_qco_to_qc.cpp +++ b/mlir/unittests/Conversion/QCOToQC/test_qco_to_qc.cpp @@ -14,8 +14,10 @@ #include "mlir/Dialect/MQT/IR/MQTDialect.h" #include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" #include "mlir/Dialect/QC/IR/QCDialect.h" +#include "mlir/Dialect/QC/IR/QCOps.h" #include "mlir/Dialect/QCO/Builder/QCOProgramBuilder.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" +#include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" #include "mlir/Support/Passes.h" #include "qc_programs.h" @@ -29,6 +31,7 @@ #include #include #include +#include #include #include #include @@ -38,6 +41,7 @@ #include #include +#include #include #include #include @@ -82,10 +86,217 @@ class QCOToQCTest : public testing::TestWithParam { } // namespace -static LogicalResult runQCOToQCConversion(ModuleOp module) { - PassManager pm(module.getContext()); +static LogicalResult runQCOToQCConversion(ModuleOp moduleOp) { + PassManager pm(moduleOp.getContext()); pm.addPass(createQCOToQC()); - return pm.run(module); + return pm.run(moduleOp); +} + +TEST(QCOToQCRegressionTest, StripsPositionalQubitResultsFromUnitaryCalls) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func private @flip(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.x %q : !qco.qubit -> !qco.qubit + return %out : !qco.qubit + } + func.func @main(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.entry_point} { + %out = qco.call @flip(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } +} +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + for (auto function : moduleOp->getOps()) { + EXPECT_EQ(function.getNumResults(), 0U); + } + std::size_t calls = 0; + moduleOp->walk([&](qc::CallOp) { ++calls; }); + EXPECT_EQ(calls, 1U); +} + +TEST(QCOToQCRegressionTest, StripsPositionalQubitResultsFromGenericCalls) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func private @reset(%q: !qco.qubit) -> (i1, !qco.qubit) { + %out = qco.reset %q : !qco.qubit -> !qco.qubit + %flag = arith.constant true + return %flag, %out : i1, !qco.qubit + } + func.func @main(%q: !qco.qubit) -> (i1, !qco.qubit) + attributes {mqt.entry_point} { + %flag, %out = func.call @reset(%q) + : (!qco.qubit) -> (i1, !qco.qubit) + return %flag, %out : i1, !qco.qubit + } +} +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + for (auto function : moduleOp->getOps()) { + ASSERT_EQ(function.getNumResults(), 1U); + EXPECT_TRUE(function.getResultTypes().front().isInteger(1)); + } + auto main = mlir::mqt::getEntryPoint(*moduleOp); + ASSERT_TRUE(main); + auto call = *main.getBody().getOps().begin(); + ASSERT_EQ(call.getNumResults(), 1U); + EXPECT_TRUE(call.getResult(0).getType().isInteger(1)); +} + +TEST(QCOToQCRegressionTest, RejectsUnrepresentableCallResultAttributes) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + constexpr auto sources = std::to_array({ + R"mlir( +module { + func.func private @reset(%q: !qco.qubit) -> (i1, !qco.qubit) { + %out = qco.reset %q : !qco.qubit -> !qco.qubit + %flag = arith.constant true + return %flag, %out : i1, !qco.qubit + } + func.func @main(%q: !qco.qubit) -> (i1, !qco.qubit) + attributes {mqt.entry_point} { + %flag, %out = func.call @reset(%q) { + res_attrs = [{}, {tag = "wire"}] + } : (!qco.qubit) -> (i1, !qco.qubit) + return %flag, %out : i1, !qco.qubit + } +} +)mlir", + R"mlir( +module { + func.func private @flip(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.x %q : !qco.qubit -> !qco.qubit + return %out : !qco.qubit + } + func.func @main(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.entry_point} { + %out = qco.call @flip(%q) {res_attrs = [{tag = "wire"}]} + : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } +} +)mlir", + }); + + for (const auto source : sources) { + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= + StringRef(diagnostic.str()).contains("cannot preserve"); + return success(); + }); + EXPECT_TRUE(failed(runQCOToQCConversion(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); + } +} + +TEST(QCOToQCRegressionTest, RejectsUnrepresentableFunctionResultAttributes) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + qco::QCOProgramBuilder builder(&context); + builder.initialize(); + auto function = builder.createFunction( + "passthrough", TypeRange{qco::QubitType::get(&context)}, + [](ValueRange arguments) { + return SmallVector{arguments.front()}; + }); + function.setResultAttr(0, "test.tag", StringAttr::get(&context, "wire")); + auto moduleOp = builder.finalize(); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= + StringRef(diagnostic.str()).contains("cannot preserve attributes"); + return success(); + }); + EXPECT_TRUE(failed(runQCOToQCConversion(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + +TEST(QCOToQCRegressionTest, RejectsNonPositionalQubitResults) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func @swap(%left: !qco.qubit, %right: !qco.qubit) + -> (!qco.qubit, !qco.qubit) { + return %right, %left : !qco.qubit, !qco.qubit + } +} +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= StringRef(diagnostic.str()) + .contains("must return its qubit arguments " + "positionally"); + return success(); + }); + EXPECT_TRUE(failed(runQCOToQCConversion(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + +TEST(QCOToQCRegressionTest, RejectsMissingPositionalQubitResults) { + DialectRegistry registry; + registry.insert(); + MLIRContext context(registry); + context.loadAllAvailableDialects(); + + auto moduleOp = parseSourceString(R"mlir(module { + func.func @bad(%q: !qco.qubit) -> i1 { + %flag = arith.constant true + return %flag : i1 + } + })mlir", + &context); + ASSERT_TRUE(moduleOp); + EXPECT_TRUE(failed(runQCOToQCConversion(*moduleOp))); } TEST(QCOToQCRegressionTest, PreservesDynamicQTensorSlotSwapAcrossLoop) { @@ -123,13 +334,13 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOToQCConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); - auto function = *module->getOps().begin(); + auto function = *moduleOp->getOps().begin(); auto loops = llvm::to_vector(function.getBody().getOps()); ASSERT_EQ(loops.size(), 1U); EXPECT_EQ(llvm::range_size(loops[0].getBody()->getOps()), @@ -252,25 +463,25 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOToQCConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::IfOp ifOp; - module->walk([&](scf::IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](scf::IfOp candidate) { ifOp = candidate; }); ASSERT_TRUE(ifOp); ASSERT_EQ(ifOp.getNumResults(), 1); EXPECT_TRUE(ifOp.getResult(0).getType().isInteger(64)); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), ifOp.getResult(0)); bool containsQCOOperations = false; - module->walk([&](Operation* operation) { + moduleOp->walk([&](Operation* operation) { containsQCOOperations |= operation->getDialect() == context.getLoadedDialect(); }); @@ -307,26 +518,26 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); - ASSERT_TRUE(succeeded(runQCOToQCConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::IndexSwitchOp switchOp; - module->walk([&](scf::IndexSwitchOp candidate) { switchOp = candidate; }); + moduleOp->walk([&](scf::IndexSwitchOp candidate) { switchOp = candidate; }); ASSERT_TRUE(switchOp); ASSERT_EQ(switchOp.getNumResults(), 1); EXPECT_TRUE(switchOp.getResult(0).getType().isInteger(64)); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), switchOp.getResult(0)); bool containsQCOOperations = false; - module->walk([&](Operation* operation) { + moduleOp->walk([&](Operation* operation) { containsQCOOperations |= operation->getDialect() == context.getLoadedDialect(); }); @@ -362,14 +573,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOToQCConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::ForOp loop; - module->walk([&](scf::ForOp candidate) { loop = candidate; }); + moduleOp->walk([&](scf::ForOp candidate) { loop = candidate; }); ASSERT_TRUE(loop); ASSERT_EQ(loop.getInitArgs().size(), 1); EXPECT_TRUE(loop.getInitArgs().front().getType().isInteger(64)); @@ -411,14 +622,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOToQCConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::WhileOp loop; - module->walk([&](scf::WhileOp candidate) { loop = candidate; }); + moduleOp->walk([&](scf::WhileOp candidate) { loop = candidate; }); ASSERT_TRUE(loop); ASSERT_EQ(loop.getInits().size(), 1); EXPECT_TRUE(loop.getInits().front().getType().isF32()); @@ -469,8 +680,7 @@ TEST_P(QCOToQCTest, ProgramEquivalence) { areModulesEquivalentWithPermutations(program.get(), reference.get())); } -/// \name QCOToQC/QubitManagement/QubitManagement.cpp -/// @{ +// QCOToQC/QubitManagement/QubitManagement.cpp INSTANTIATE_TEST_SUITE_P( QCOQubitManagementTest, QCOToQCTest, testing::Values( @@ -494,10 +704,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"AllocDeallocPair", MQT_NAMED_BUILDER(qco::allocSinkPair), MQT_NAMED_BUILDER(qc::allocDeallocPair)})); -/// @} -/// \name QCOToQC/Modifiers/PowOp.cpp -/// @{ +// QCOToQC/Modifiers/PowOp.cpp INSTANTIATE_TEST_SUITE_P( QCOPowOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"CtrlPowSx", @@ -506,10 +714,7 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"PowTwo", MQT_NAMED_BUILDER(qco::powTwo), MQT_NAMED_BUILDER(qc::powTwo)})); -/// @} - -/// \name QCOToQC/Modifiers/CtrlOp.cpp -/// @{ +// QCOToQC/Modifiers/CtrlOp.cpp INSTANTIATE_TEST_SUITE_P( QCOCtrlOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"CtrlTwo", MQT_NAMED_BUILDER(qco::ctrlTwo), @@ -520,10 +725,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"CtrlInvTwo", MQT_NAMED_BUILDER(qco::ctrlInvTwo), MQT_NAMED_BUILDER(qc::ctrlInvTwo)})); -/// @} -/// \name QCOToQC/Modifiers/InvOp.cpp -/// @{ +// QCOToQC/Modifiers/InvOp.cpp INSTANTIATE_TEST_SUITE_P( QCOInvOpTest, QCOToQCTest, testing::Values( @@ -541,10 +744,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qc::multipleControlledDcx)}, QCOToQCTestCase{"InvTwo", MQT_NAMED_BUILDER(qco::invTwo), MQT_NAMED_BUILDER(qc::invTwo)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/BarrierOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/BarrierOp.cpp INSTANTIATE_TEST_SUITE_P( QCOBarrierOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"Barrier", MQT_NAMED_BUILDER(qco::barrier), @@ -556,10 +757,8 @@ INSTANTIATE_TEST_SUITE_P( "BarrierMultipleQubits", MQT_NAMED_BUILDER(qco::barrierMultipleQubits), MQT_NAMED_BUILDER(qc::barrierMultipleQubits)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/DcxOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/DcxOp.cpp INSTANTIATE_TEST_SUITE_P( QCODCXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"DCX", MQT_NAMED_BUILDER(qco::dcx), @@ -571,10 +770,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledDCX", MQT_NAMED_BUILDER(qco::multipleControlledDcx), MQT_NAMED_BUILDER(qc::multipleControlledDcx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/EcrOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/EcrOp.cpp INSTANTIATE_TEST_SUITE_P( QCOECROpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"ECR", MQT_NAMED_BUILDER(qco::ecr), @@ -586,18 +783,14 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledECR", MQT_NAMED_BUILDER(qco::multipleControlledEcr), MQT_NAMED_BUILDER(qc::multipleControlledEcr)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/GphaseOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/GphaseOp.cpp INSTANTIATE_TEST_SUITE_P(QCOGPhaseOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{ "GlobalPhase", MQT_NAMED_BUILDER(qco::globalPhase), MQT_NAMED_BUILDER(qc::globalPhase)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/HOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/HOp.cpp INSTANTIATE_TEST_SUITE_P( QCOHOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"H", MQT_NAMED_BUILDER(qco::h), @@ -611,10 +804,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"HWithoutRegister", MQT_NAMED_BUILDER(qco::hWithoutRegister), MQT_NAMED_BUILDER(qc::hWithoutRegister)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/IswapOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/IswapOp.cpp INSTANTIATE_TEST_SUITE_P( QCOiSWAPOpTest, QCOToQCTest, testing::Values( @@ -626,10 +817,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControllediSWAP", MQT_NAMED_BUILDER(qco::multipleControlledIswap), MQT_NAMED_BUILDER(qc::multipleControlledIswap)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/POp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/POp.cpp INSTANTIATE_TEST_SUITE_P( QCOPOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"P", MQT_NAMED_BUILDER(qco::p), @@ -641,10 +830,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledP", MQT_NAMED_BUILDER(qco::multipleControlledP), MQT_NAMED_BUILDER(qc::multipleControlledP)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RCCXOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RCCXOp.cpp INSTANTIATE_TEST_SUITE_P( QCORCCXOpTest, QCOToQCTest, testing::Values( @@ -656,10 +843,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControlledRCCX", MQT_NAMED_BUILDER(qco::multipleControlledRccx), MQT_NAMED_BUILDER(qc::multipleControlledRccx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/ROp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/ROp.cpp INSTANTIATE_TEST_SUITE_P( QCOROpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"R", MQT_NAMED_BUILDER(qco::r), @@ -671,10 +856,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledR", MQT_NAMED_BUILDER(qco::multipleControlledR), MQT_NAMED_BUILDER(qc::multipleControlledR)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RxOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RX", MQT_NAMED_BUILDER(qco::rx), @@ -686,10 +869,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRX", MQT_NAMED_BUILDER(qco::multipleControlledRx), MQT_NAMED_BUILDER(qc::multipleControlledRx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RxxOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RxxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORXXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RXX", MQT_NAMED_BUILDER(qco::rxx), @@ -701,10 +882,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRXX", MQT_NAMED_BUILDER(qco::multipleControlledRxx), MQT_NAMED_BUILDER(qc::multipleControlledRxx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RyOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RyOp.cpp INSTANTIATE_TEST_SUITE_P( QCORYOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RY", MQT_NAMED_BUILDER(qco::ry), @@ -716,10 +895,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRY", MQT_NAMED_BUILDER(qco::multipleControlledRy), MQT_NAMED_BUILDER(qc::multipleControlledRy)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RyyOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RyyOp.cpp INSTANTIATE_TEST_SUITE_P( QCORYYOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RYY", MQT_NAMED_BUILDER(qco::ryy), @@ -731,10 +908,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRYY", MQT_NAMED_BUILDER(qco::multipleControlledRyy), MQT_NAMED_BUILDER(qc::multipleControlledRyy)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RzOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RzOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RZ", MQT_NAMED_BUILDER(qco::rz), @@ -746,10 +921,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRZ", MQT_NAMED_BUILDER(qco::multipleControlledRz), MQT_NAMED_BUILDER(qc::multipleControlledRz)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RzxOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RzxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RZX", MQT_NAMED_BUILDER(qco::rzx), @@ -761,10 +934,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRZX", MQT_NAMED_BUILDER(qco::multipleControlledRzx), MQT_NAMED_BUILDER(qc::multipleControlledRzx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/RzzOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/RzzOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZZOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"RZZ", MQT_NAMED_BUILDER(qco::rzz), @@ -776,10 +947,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRZZ", MQT_NAMED_BUILDER(qco::multipleControlledRzz), MQT_NAMED_BUILDER(qc::multipleControlledRzz)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/SOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/SOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"S", MQT_NAMED_BUILDER(qco::s), @@ -791,10 +960,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledS", MQT_NAMED_BUILDER(qco::multipleControlledS), MQT_NAMED_BUILDER(qc::multipleControlledS)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/SdgOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/SdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSdgOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"Sdg", MQT_NAMED_BUILDER(qco::sdg), @@ -806,10 +973,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledSdg", MQT_NAMED_BUILDER(qco::multipleControlledSdg), MQT_NAMED_BUILDER(qc::multipleControlledSdg)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/SwapOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/SwapOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSWAPOpTest, QCOToQCTest, testing::Values( @@ -821,10 +986,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControlledSWAP", MQT_NAMED_BUILDER(qco::multipleControlledSwap), MQT_NAMED_BUILDER(qc::multipleControlledSwap)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/SxOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/SxOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"SX", MQT_NAMED_BUILDER(qco::sx), @@ -836,10 +999,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledSX", MQT_NAMED_BUILDER(qco::multipleControlledSx), MQT_NAMED_BUILDER(qc::multipleControlledSx)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/SxdgOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/SxdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSXdgOpTest, QCOToQCTest, testing::Values( @@ -851,10 +1012,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControlledSXdg", MQT_NAMED_BUILDER(qco::multipleControlledSxdg), MQT_NAMED_BUILDER(qc::multipleControlledSxdg)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/TOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/TOp.cpp INSTANTIATE_TEST_SUITE_P( QCOTOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"T", MQT_NAMED_BUILDER(qco::t_), @@ -866,10 +1025,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledT", MQT_NAMED_BUILDER(qco::multipleControlledT), MQT_NAMED_BUILDER(qc::multipleControlledT)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/TdgOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/TdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOTdgOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"Tdg", MQT_NAMED_BUILDER(qco::tdg), @@ -881,10 +1038,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledTdg", MQT_NAMED_BUILDER(qco::multipleControlledTdg), MQT_NAMED_BUILDER(qc::multipleControlledTdg)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/U2Op.cpp -/// @{ +// QCOToQC/Operations/StandardGates/U2Op.cpp INSTANTIATE_TEST_SUITE_P( QCOU2OpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"U2", MQT_NAMED_BUILDER(qco::u2), @@ -896,10 +1051,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledU2", MQT_NAMED_BUILDER(qco::multipleControlledU2), MQT_NAMED_BUILDER(qc::multipleControlledU2)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/UOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/UOp.cpp INSTANTIATE_TEST_SUITE_P( QCOUOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"U", MQT_NAMED_BUILDER(qco::u), @@ -911,10 +1064,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledU", MQT_NAMED_BUILDER(qco::multipleControlledU), MQT_NAMED_BUILDER(qc::multipleControlledU)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/XOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/XOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"X", MQT_NAMED_BUILDER(qco::x), @@ -929,10 +1080,8 @@ INSTANTIATE_TEST_SUITE_P( "RepeatedControlledX", MQT_NAMED_BUILDER(qco::repeatedControlledX), MQT_NAMED_BUILDER(qc::repeatedControlledX)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/XxMinusYyOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/XxMinusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXXMinusYYOpTest, QCOToQCTest, testing::Values( @@ -944,10 +1093,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControlledXXMinusYY", MQT_NAMED_BUILDER(qco::multipleControlledXxMinusYY), MQT_NAMED_BUILDER(qc::multipleControlledXxMinusYY)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/XxPlusYyOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/XxPlusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXXPlusYYOpTest, QCOToQCTest, testing::Values( @@ -959,10 +1106,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MultipleControlledXXPlusYY", MQT_NAMED_BUILDER(qco::multipleControlledXxPlusYY), MQT_NAMED_BUILDER(qc::multipleControlledXxPlusYY)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/YOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/YOp.cpp INSTANTIATE_TEST_SUITE_P( QCOYOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"Y", MQT_NAMED_BUILDER(qco::y), @@ -974,10 +1119,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledY", MQT_NAMED_BUILDER(qco::multipleControlledY), MQT_NAMED_BUILDER(qc::multipleControlledY)})); -/// @} -/// \name QCOToQC/Operations/StandardGates/ZOp.cpp -/// @{ +// QCOToQC/Operations/StandardGates/ZOp.cpp INSTANTIATE_TEST_SUITE_P( QCOZOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"Z", MQT_NAMED_BUILDER(qco::z), @@ -989,10 +1132,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledZ", MQT_NAMED_BUILDER(qco::multipleControlledZ), MQT_NAMED_BUILDER(qc::multipleControlledZ)})); -/// @} -/// \name QCOToQC/Operations/MeasureOp.cpp -/// @{ +// QCOToQC/Operations/MeasureOp.cpp INSTANTIATE_TEST_SUITE_P( QCOMeasureOpTest, QCOToQCTest, testing::Values( @@ -1019,10 +1160,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"MeasurementWithoutRegisters", MQT_NAMED_BUILDER(qco::measurementWithoutRegisters), MQT_NAMED_BUILDER(qc::measurementWithoutRegisters)})); -/// @} -/// \name QCOToQC/Operations/ResetOp.cpp -/// @{ +// QCOToQC/Operations/ResetOp.cpp INSTANTIATE_TEST_SUITE_P( QCOResetOpTest, QCOToQCTest, testing::Values( @@ -1036,10 +1175,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"RepeatedResetAfterSingleOp", MQT_NAMED_BUILDER(qco::repeatedResetAfterSingleOp), MQT_NAMED_BUILDER(qc::resetQubitAfterSingleOp)})); -/// @} -/// \name QCOToQC/Operations/IfOp.cpp -/// @{ +// QCOToQC/Operations/IfOp.cpp INSTANTIATE_TEST_SUITE_P( QCOIfOpTest, QCOToQCTest, testing::Values( @@ -1057,10 +1194,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"NestedIfOpForLoop", MQT_NAMED_BUILDER(qco::nestedIfOpForLoop), MQT_NAMED_BUILDER(qc::nestedIfOpForLoop)})); -/// @} -/// \name QCOToQC/Operations/IndexSwitchOp.cpp -/// @{ +// QCOToQC/Operations/IndexSwitchOp.cpp INSTANTIATE_TEST_SUITE_P( QCOIndexSwitchOpTest, QCOToQCTest, testing::Values(QCOToQCTestCase{"SimpleIndexSwitchOp", @@ -1070,10 +1205,8 @@ INSTANTIATE_TEST_SUITE_P( "IndexSwitchMultiCase", MQT_NAMED_BUILDER(qco::indexSwitchMultiCase), MQT_NAMED_BUILDER(qc::indexSwitchMultiCase)})); -/// @} -/// \name QCOToQC/Operations/WhileOp.cpp -/// @{ +// QCOToQC/Operations/WhileOp.cpp INSTANTIATE_TEST_SUITE_P( SCFWhileOpTest, QCOToQCTest, testing::Values( @@ -1082,10 +1215,8 @@ INSTANTIATE_TEST_SUITE_P( QCOToQCTestCase{"SimpleDoWhile", MQT_NAMED_BUILDER(qco::simpleDoWhileReset), MQT_NAMED_BUILDER(qc::simpleDoWhileReset)})); -/// @} -/// \name QCOToQC/Operations/ForOp.cpp -/// @{ +// QCOToQC/Operations/ForOp.cpp INSTANTIATE_TEST_SUITE_P( SCFForOpTest, QCOToQCTest, testing::Values( @@ -1109,4 +1240,3 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qco::nestedForLoopCtrlOpWithExtractedQubit), MQT_NAMED_BUILDER( aliasSafeNestedForLoopCtrlOpWithExtractedQubit)})); -/// @} diff --git a/mlir/unittests/Conversion/QCQCORoundTrip/CMakeLists.txt b/mlir/unittests/Conversion/QCQCORoundTrip/CMakeLists.txt index 62a3f61d0c..740a35cc75 100644 --- a/mlir/unittests/Conversion/QCQCORoundTrip/CMakeLists.txt +++ b/mlir/unittests/Conversion/QCQCORoundTrip/CMakeLists.txt @@ -19,7 +19,11 @@ target_link_libraries( MLIRMQTDialect MLIRParser MLIRPass + MLIRQCPrograms + MLIRQCProgramBuilder MLIRQCDialect + MLIRQCOPrograms + MLIRQCOProgramBuilder MLIRQCODialect MLIRSCFDialect MLIRSupportMQT diff --git a/mlir/unittests/Conversion/QCQCORoundTrip/test_qc_qco_round_trip.cpp b/mlir/unittests/Conversion/QCQCORoundTrip/test_qc_qco_round_trip.cpp index cda3572613..d5763dc791 100644 --- a/mlir/unittests/Conversion/QCQCORoundTrip/test_qc_qco_round_trip.cpp +++ b/mlir/unittests/Conversion/QCQCORoundTrip/test_qc_qco_round_trip.cpp @@ -8,16 +8,22 @@ * Licensed under the MIT License */ +#include "TestCaseUtils.h" #include "mlir/Conversion/QCOToQC/QCOToQC.h" #include "mlir/Conversion/QCToQCO/QCToQCO.h" #include "mlir/Dialect/CBit/IR/CBitAttributes.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" #include "mlir/Dialect/CBit/IR/CBitOps.h" #include "mlir/Dialect/MQT/IR/MQTDialect.h" +#include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" #include "mlir/Dialect/QC/IR/QCDialect.h" #include "mlir/Dialect/QC/IR/QCOps.h" +#include "mlir/Dialect/QCO/Builder/QCOProgramBuilder.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" +#include "mlir/Dialect/QCO/IR/QCOOps.h" #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" +#include "qc_programs.h" +#include "qco_programs.h" #include #include @@ -36,7 +42,9 @@ #include #include +#include #include +#include using namespace mlir; @@ -49,23 +57,30 @@ class QCQCORoundTripTest : public testing::Test { QCQCORoundTripTest() { DialectRegistry registry; registry - .insert(); context.appendDialectRegistry(registry); context.loadAllAvailableDialects(); } - [[nodiscard]] LogicalResult runRoundTrip(ModuleOp module) { + [[nodiscard]] LogicalResult runRoundTrip(ModuleOp moduleOp) { PassManager pm(&context); pm.addPass(createQCToQCO()); pm.addPass(createQCOToQC()); - return pm.run(module); + return pm.run(moduleOp); } - static void expectNoScratchStorage(ModuleOp module) { + [[nodiscard]] LogicalResult runReverseRoundTrip(ModuleOp moduleOp) { + PassManager pm(&context); + pm.addPass(createQCOToQC()); + pm.addPass(createQCToQCO()); + return pm.run(moduleOp); + } + + static void expectNoScratchStorage(ModuleOp moduleOp) { bool containsScratchStorage = false; - module.walk([&](Operation* operation) { + moduleOp.walk([&](Operation* operation) { containsScratchStorage |= isa(operation); }); @@ -79,7 +94,7 @@ TEST_F(QCQCORoundTripTest, PreservesSharedMQTMetadata) { constexpr StringLiteral source = R"mlir( module { func.func @main(%theta: f64 {mqt.input_name = "theta"}) - attributes {mqt.entry_point} { + attributes {mqt.entry_point, mqt.source_name = "source"} { %reg = memref.alloc() {mqt.register_name = "q"} : memref<2x!qc.qubit> memref.dealloc %reg : memref<2x!qc.qubit> @@ -95,9 +110,13 @@ module { auto function = moduleOp->lookupSymbol("main"); ASSERT_TRUE(function); - EXPECT_TRUE(mqt::isEntryPoint(function)); + EXPECT_TRUE(::mlir::mqt::isEntryPoint(function)); + auto sourceName = function->getAttrOfType( + ::mlir::mqt::MQTDialect::SourceNameAttrHelper::getNameStr()); + ASSERT_TRUE(sourceName); + EXPECT_EQ(sourceName.getValue(), "source"); const auto inputName = function.getArgAttrOfType( - 0, mqt::MQTDialect::InputNameAttrHelper::getNameStr()); + 0, ::mlir::mqt::MQTDialect::InputNameAttrHelper::getNameStr()); ASSERT_TRUE(inputName); EXPECT_EQ(inputName.getValue(), "theta"); @@ -105,11 +124,41 @@ module { moduleOp->walk([&](memref::AllocOp op) { allocation = op; }); ASSERT_TRUE(allocation); const auto registerName = allocation->getAttrOfType( - mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()); + ::mlir::mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()); ASSERT_TRUE(registerName); EXPECT_EQ(registerName.getValue(), "q"); } +TEST_F(QCQCORoundTripTest, PreservesReusableFunctions) { + const std::array cases{ + std::tuple{MQT_NAMED_BUILDER(qc::reusableUnitaryFunction), + MQT_NAMED_BUILDER(qco::reusableUnitaryFunction), true}, + std::tuple{MQT_NAMED_BUILDER(qc::reusableResetFunction), + MQT_NAMED_BUILDER(qco::reusableResetFunction), false}}; + + for (const auto& [qcBuilder, qcoBuilder, unitary] : cases) { + SCOPED_TRACE(qcBuilder.name); + auto qcModule = ::mqt::test::buildMLIRProgram(&context, qcBuilder); + auto qcoModule = ::mqt::test::buildMLIRProgram(&context, qcoBuilder); + ASSERT_TRUE(qcModule); + ASSERT_TRUE(qcoModule); + ASSERT_TRUE(succeeded(runRoundTrip(*qcModule))); + ASSERT_TRUE(succeeded(runReverseRoundTrip(*qcoModule))); + ASSERT_TRUE(succeeded(verify(*qcModule))); + ASSERT_TRUE(succeeded(verify(*qcoModule))); + + for (ModuleOp moduleOp : {*qcModule, *qcoModule}) { + size_t unitaryCalls = 0; + size_t genericCalls = 0; + moduleOp.walk([&](qc::CallOp) { ++unitaryCalls; }); + moduleOp.walk([&](qco::CallOp) { ++unitaryCalls; }); + moduleOp.walk([&](func::CallOp) { ++genericCalls; }); + EXPECT_EQ(unitaryCalls, unitary ? 1U : 0U); + EXPECT_EQ(genericCalls, unitary ? 0U : 1U); + } + } +} + TEST_F(QCQCORoundTripTest, PreservesClassicalRegistersWithoutConversion) { constexpr llvm::StringLiteral source = R"mlir( module { @@ -144,18 +193,20 @@ module { ASSERT_EQ(loads.size(), 1); ASSERT_EQ(stores.size(), 2); EXPECT_EQ(allocations[0].getInitialization(), cbit::Initialization::Zero); - EXPECT_EQ(allocations[0] - ->getAttrOfType( - mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()) - .getValue(), - "zero"); + EXPECT_EQ( + allocations[0] + ->getAttrOfType( + ::mlir::mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()) + .getValue(), + "zero"); EXPECT_EQ(allocations[1].getInitialization(), cbit::Initialization::Undefined); - EXPECT_EQ(allocations[1] - ->getAttrOfType( - mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()) - .getValue(), - "undefined"); + EXPECT_EQ( + allocations[1] + ->getAttrOfType( + ::mlir::mqt::MQTDialect::RegisterNameAttrHelper::getNameStr()) + .getValue(), + "undefined"); EXPECT_EQ(loads.front().getReg(), allocations.front().getResult()); EXPECT_EQ(stores.front().getReg(), allocations.front().getResult()); EXPECT_EQ(stores.back().getReg(), allocations.back().getResult()); @@ -187,22 +238,22 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runRoundTrip(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runRoundTrip(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::IfOp ifOp; - module->walk([&](scf::IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](scf::IfOp candidate) { ifOp = candidate; }); ASSERT_TRUE(ifOp); ASSERT_EQ(ifOp.getNumResults(), 1); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), ifOp.getResult(0)); - expectNoScratchStorage(*module); + expectNoScratchStorage(*moduleOp); } TEST_F(QCQCORoundTripTest, PreservesClassicalIndexSwitchResultWithoutScratch) { @@ -228,22 +279,22 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runRoundTrip(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runRoundTrip(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::IndexSwitchOp switchOp; - module->walk([&](scf::IndexSwitchOp candidate) { switchOp = candidate; }); + moduleOp->walk([&](scf::IndexSwitchOp candidate) { switchOp = candidate; }); ASSERT_TRUE(switchOp); ASSERT_EQ(switchOp.getNumResults(), 1); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), switchOp.getResult(0)); - expectNoScratchStorage(*module); + expectNoScratchStorage(*moduleOp); } TEST_F(QCQCORoundTripTest, PreservesDenseUnitaryMatrixAndQubitArity) { @@ -265,18 +316,18 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); ElementsAttr originalMatrix; - module->walk( + moduleOp->walk( [&](qc::UnitaryOp unitary) { originalMatrix = unitary.getMatrix(); }); ASSERT_TRUE(originalMatrix); std::string serialized; llvm::raw_string_ostream stream(serialized); - module->print(stream); + moduleOp->print(stream); stream.flush(); auto reparsed = parseSourceString(serialized, &context); ASSERT_TRUE(reparsed); @@ -287,11 +338,11 @@ module { EXPECT_EQ(reparsedUnitary.getQubits().size(), 2U); EXPECT_EQ(reparsedUnitary.getMatrix(), originalMatrix); - ASSERT_TRUE(succeeded(runRoundTrip(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + ASSERT_TRUE(succeeded(runRoundTrip(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); qc::UnitaryOp unitary; - module->walk([&](qc::UnitaryOp candidate) { unitary = candidate; }); + moduleOp->walk([&](qc::UnitaryOp candidate) { unitary = candidate; }); ASSERT_TRUE(unitary); EXPECT_EQ(unitary.getQubits().size(), 2U); EXPECT_EQ(unitary.getMatrix(), originalMatrix); diff --git a/mlir/unittests/Conversion/QCToQCO/CMakeLists.txt b/mlir/unittests/Conversion/QCToQCO/CMakeLists.txt index b595155600..1bcb89dc62 100644 --- a/mlir/unittests/Conversion/QCToQCO/CMakeLists.txt +++ b/mlir/unittests/Conversion/QCToQCO/CMakeLists.txt @@ -18,7 +18,8 @@ target_link_libraries( MLIRQCOProgramBuilder MLIRQCPrograms MLIRQCOPrograms - MLIRQCToQCO) + MLIRQCToQCO + MLIRQCOToQC) mqt_mlir_configure_unittest_target(${target_name}) diff --git a/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp b/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp index 31f4783e52..6dd3cadf31 100644 --- a/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp +++ b/mlir/unittests/Conversion/QCToQCO/test_qc_to_qco.cpp @@ -11,6 +11,7 @@ #include "Support/IRVerification.h" #include "TestCaseUtils.h" #include "mlir/Conversion/ConversionUtils.h" +#include "mlir/Conversion/QCOToQC/QCOToQC.h" #include "mlir/Conversion/QCToQCO/QCToQCO.h" #include "mlir/Dialect/CBit/IR/CBitAttributes.h" #include "mlir/Dialect/CBit/IR/CBitDialect.h" @@ -18,6 +19,7 @@ #include "mlir/Dialect/MQT/IR/MQTDialect.h" #include "mlir/Dialect/QC/Builder/QCProgramBuilder.h" #include "mlir/Dialect/QC/IR/QCDialect.h" +#include "mlir/Dialect/QC/IR/QCOps.h" #include "mlir/Dialect/QCO/Builder/QCOProgramBuilder.h" #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/IR/QCOInterfaces.h" @@ -46,6 +48,7 @@ #include #include #include +#include #include #include #include @@ -106,10 +109,16 @@ class QCToQCOTest : public testing::TestWithParam { } // namespace -static LogicalResult runQCToQCOConversion(ModuleOp module) { - PassManager pm(module.getContext()); +static LogicalResult runQCToQCOConversion(ModuleOp moduleOp) { + PassManager pm(moduleOp.getContext()); pm.addPass(createQCToQCO()); - return pm.run(module); + return pm.run(moduleOp); +} + +static LogicalResult runQCOToQCConversion(ModuleOp moduleOp) { + PassManager pm(moduleOp.getContext()); + pm.addPass(createQCOToQC()); + return pm.run(moduleOp); } namespace { @@ -128,9 +137,9 @@ class QCToQCORegressionTest : public testing::Test { context.loadAllAvailableDialects(); } - void expectNoQCOperations(ModuleOp module) { + void expectNoQCOperations(ModuleOp moduleOp) { bool retainsQCOperations = false; - module.walk([&](Operation* operation) { + moduleOp.walk([&](Operation* operation) { retainsQCOperations |= operation->getDialect() == context.getLoadedDialect(); }); @@ -237,8 +246,8 @@ class RejectingRegionMovePattern final if (!op->hasAttr("test.reject_region_move")) { return failure(); } - auto module = op->getParentOfType(); - auto destination = module.lookupSymbol("destination"); + auto moduleOp = op->getParentOfType(); + auto destination = moduleOp.lookupSymbol("destination"); if (!destination) { return failure(); } @@ -267,9 +276,9 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); TypeConverter typeConverter; typeConverter.addConversion([](Type type) -> std::optional { @@ -290,11 +299,11 @@ module { ScopedDiagnosticHandler handler( &context, [](Diagnostic& /*diagnostic*/) { return success(); }); EXPECT_TRUE( - failed(applyPartialConversion(*module, target, std::move(patterns)))); + failed(applyPartialConversion(*moduleOp, target, std::move(patterns)))); EXPECT_TRUE(sourcePreserved); - auto sourceFunc = module->lookupSymbol("source"); - auto destination = module->lookupSymbol("destination"); + auto sourceFunc = moduleOp->lookupSymbol("source"); + auto destination = moduleOp->lookupSymbol("destination"); ASSERT_TRUE(sourceFunc); ASSERT_TRUE(destination); EXPECT_FALSE(sourceFunc.getBody().empty()); @@ -320,14 +329,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); bool sawLoop = false; - module->walk([&](scf::ForOp loop) { + moduleOp->walk([&](scf::ForOp loop) { sawLoop = true; EXPECT_EQ(loop.getNumResults(), 2); EXPECT_TRUE(loop.getResult(0).getType().isInteger(1)); @@ -338,7 +347,7 @@ module { }); EXPECT_TRUE(sawLoop); - expectNoQCOperations(*module); + expectNoQCOperations(*moduleOp); } TEST_F(QCToQCORegressionTest, CoalescesStaticQubitsAcrossRegions) { @@ -408,13 +417,13 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); bool sawWhile = false; - module->walk([&](scf::WhileOp loop) { + moduleOp->walk([&](scf::WhileOp loop) { sawWhile = true; ASSERT_EQ(loop.getNumResults(), 3); EXPECT_TRUE(loop.getResult(0).getType().isInteger(64)); @@ -429,9 +438,9 @@ module { llvm::equal(yield.getOperandTypes(), loop.getInits().getTypes())); }); EXPECT_TRUE(sawWhile); - expectNoQCOperations(*module); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(*module))); - auto main = module->lookupSymbol("main"); + expectNoQCOperations(*moduleOp); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(*moduleOp))); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); APInt result; @@ -465,15 +474,15 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); - expectNoQCOperations(*module); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + expectNoQCOperations(*moduleOp); bool retainsClassicalRegister = false; - module->walk([&](memref::LoadOp op) { + moduleOp->walk([&](memref::LoadOp op) { retainsClassicalRegister |= op.getMemRefType().getElementType().isInteger(1); }); @@ -503,14 +512,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); scf::WhileOp loop; - module->walk([&](scf::WhileOp candidate) { loop = candidate; }); + moduleOp->walk([&](scf::WhileOp candidate) { loop = candidate; }); ASSERT_TRUE(loop); ASSERT_EQ(loop.getInits().size(), 2); EXPECT_TRUE(loop.getInits().front().getType().isF32()); @@ -525,7 +534,7 @@ module { llvm::equal(condition.getArgs().getTypes(), loop.getResultTypes())); auto yield = cast(loop.getAfterBody()->getTerminator()); EXPECT_TRUE(llvm::equal(yield.getOperandTypes(), loop.getInits().getTypes())); - expectNoQCOperations(*module); + expectNoQCOperations(*moduleOp); } TEST_F(QCToQCORegressionTest, LeavesUnrelatedSCFTerminatorsUntouched) { @@ -544,16 +553,16 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); bool sawExecuteRegion = false; - module->walk([&](scf::ExecuteRegionOp) { sawExecuteRegion = true; }); + moduleOp->walk([&](scf::ExecuteRegionOp) { sawExecuteRegion = true; }); EXPECT_TRUE(sawExecuteRegion); - expectNoQCOperations(*module); + expectNoQCOperations(*moduleOp); } TEST_F(QCToQCORegressionTest, PreservesIfClassicalResultsWithoutScratch) { @@ -577,14 +586,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); qco::IfOp ifOp; - module->walk([&](qco::IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](qco::IfOp candidate) { ifOp = candidate; }); ASSERT_TRUE(ifOp); ASSERT_EQ(ifOp.getClassicalResults().size(), 1); EXPECT_TRUE(ifOp.getClassicalResults().front().getType().isInteger(64)); @@ -596,18 +605,18 @@ module { EXPECT_TRUE(isa(yield.getOperand(1).getType())); } - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), ifOp.getClassicalResults().front()); bool containsScratchStorage = false; - module->walk([&](Operation* operation) { + moduleOp->walk([&](Operation* operation) { containsScratchStorage |= isa(operation); }); EXPECT_FALSE(containsScratchStorage); - expectNoQCOperations(*module); + expectNoQCOperations(*moduleOp); } TEST_F(QCToQCORegressionTest, @@ -634,14 +643,14 @@ module { } )mlir"; - auto module = parseSourceString(source, &context); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCToQCOConversion(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); qco::IndexSwitchOp switchOp; - module->walk([&](qco::IndexSwitchOp candidate) { switchOp = candidate; }); + moduleOp->walk([&](qco::IndexSwitchOp candidate) { switchOp = candidate; }); ASSERT_TRUE(switchOp); ASSERT_EQ(switchOp.getClassicalResults().size(), 1); EXPECT_TRUE(switchOp.getClassicalResults().front().getType().isInteger(64)); @@ -655,18 +664,18 @@ module { EXPECT_TRUE(isa(yield.getOperand(1).getType())); } - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); EXPECT_EQ(returnOp.getOperand(0), switchOp.getClassicalResults().front()); bool containsScratchStorage = false; - module->walk([&](Operation* operation) { + moduleOp->walk([&](Operation* operation) { containsScratchStorage |= isa(operation); }); EXPECT_FALSE(containsScratchStorage); - expectNoQCOperations(*module); + expectNoQCOperations(*moduleOp); } TEST_F(QCToQCORegressionTest, @@ -759,15 +768,21 @@ module { EXPECT_EQ(name.getValue(), "named_qubits"); } -TEST_F(QCToQCORegressionTest, RejectsRegisterBackedReferenceEscapes) { +TEST_F(QCToQCORegressionTest, ConvertsRegisterBackedGenericCalls) { constexpr llvm::StringLiteral source = R"mlir( module { - func.func private @escape(!qc.qubit) + func.func private @reset(%flag: i1, %q: !qc.qubit, %value: i64) -> i1 { + qc.reset %q : !qc.qubit + return %flag : i1 + } func.func @main() attributes {mqt.entry_point} { %reg = memref.alloc() : memref<1x!qc.qubit> %c0 = arith.constant 0 : index %q = memref.load %reg[%c0] : memref<1x!qc.qubit> - func.call @escape(%q) : (!qc.qubit) -> () + %true = arith.constant true + %value = arith.constant 42 : i64 + %result = func.call @reset(%true, %q, %value) + : (i1, !qc.qubit, i64) -> i1 memref.dealloc %reg : memref<1x!qc.qubit> return } @@ -777,16 +792,20 @@ module { auto moduleOp = parseSourceString(source, &context); ASSERT_TRUE(moduleOp); ASSERT_TRUE(succeeded(verify(*moduleOp))); - - bool sawExpectedDiagnostic = false; - ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { - sawExpectedDiagnostic |= - StringRef(diagnostic.str()) - .contains("cannot consume a register-backed qubit reference"); - return success(); - }); - EXPECT_TRUE(failed(runQCToQCOConversion(*moduleOp))); - EXPECT_TRUE(sawExpectedDiagnostic); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + EXPECT_TRUE(succeeded(qco::verifyLinearity(*moduleOp))); + auto call = *mlir::mqt::getEntryPoint(*moduleOp) + .getBody() + .getOps() + .begin(); + ASSERT_EQ(call.getNumOperands(), 3U); + EXPECT_TRUE(isa(call.getOperand(1).getType())); + ASSERT_EQ(call.getNumResults(), 2U); + EXPECT_TRUE(isa(call.getResult(1).getType())); + EXPECT_TRUE(call.getOperand(1).getDefiningOp()); + ASSERT_TRUE(call.getResult(1).hasOneUse()); + EXPECT_TRUE(isa(*call.getResult(1).getUsers().begin())); } TEST_F(QCToQCORegressionTest, PreflightRejectsNonOneDimensionalQubitRegisters) { @@ -852,10 +871,8 @@ module { EXPECT_TRUE(sawExpectedDiagnostic); } -TEST_F(QCToQCORegressionTest, - PreflightRejectsUnsupportedQuantumBlockArguments) { - constexpr auto sources = std::to_array({ - R"mlir( +TEST_F(QCToQCORegressionTest, ConvertsQubitFunctionArgumentsToTrailingResults) { + constexpr llvm::StringLiteral source = R"mlir( module { func.func @main(%q: !qc.qubit) attributes {mqt.entry_point} { @@ -863,7 +880,206 @@ module { return } } +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + + auto function = *moduleOp->getOps().begin(); + ASSERT_EQ(function.getNumArguments(), 1U); + EXPECT_TRUE(isa(function.getArgument(0).getType())); + ASSERT_EQ(function.getNumResults(), 1U); + EXPECT_TRUE(isa(function.getResultTypes().front())); + auto x = *function.getBody().front().getOps().begin(); + EXPECT_EQ( + cast(function.getBody().front().back()).getOperand(0), + x.getQubitOut()); +} + +TEST_F(QCToQCORegressionTest, RoundTripsUnitaryFunctionCalls) { + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func private @flip(%q: !qc.qubit) attributes {mqt.unitary} { + qc.x %q : !qc.qubit + return + } + func.func private @reset(%q: !qc.qubit) -> i1 { + qc.reset %q : !qc.qubit + %flag = arith.constant true + return %flag : i1 + } + func.func @main(%q: !qc.qubit) -> i1 attributes {mqt.entry_point} { + qc.call @flip(%q) { + arg_attrs = [{tag = "unitary-input"}], tag = "unitary-call" + } + : !qc.qubit + %flag = func.call @reset(%q) { + arg_attrs = [{tag = "generic-input"}], no_inline, + res_attrs = [{tag = "ordinary-result"}], tag = "generic-call" + } : (!qc.qubit) -> i1 + return %flag : i1 + } +} +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + SmallVector qcoCalls; + moduleOp->walk([&](qco::CallOp call) { qcoCalls.emplace_back(call); }); + ASSERT_EQ(qcoCalls.size(), 1U); + ASSERT_TRUE(qcoCalls.front().getArgAttrsAttr()); + EXPECT_EQ(cast(qcoCalls.front().getArgAttrsAttr()[0]) + .getAs("tag") + .getValue(), + "unitary-input"); + EXPECT_EQ(qcoCalls.front()->getAttrOfType("tag").getValue(), + "unitary-call"); + EXPECT_FALSE(qcoCalls.front().getResAttrsAttr()); + + SmallVector genericCalls; + moduleOp->walk([&](func::CallOp call) { genericCalls.emplace_back(call); }); + ASSERT_EQ(genericCalls.size(), 1U); + EXPECT_TRUE(genericCalls.front().getNoInline()); + EXPECT_EQ(genericCalls.front()->getAttrOfType("tag").getValue(), + "generic-call"); + ASSERT_EQ(genericCalls.front().getResAttrsAttr().size(), 2U); + EXPECT_EQ(cast(genericCalls.front().getResAttrsAttr()[0]) + .getAs("tag") + .getValue(), + "ordinary-result"); + EXPECT_TRUE( + cast(genericCalls.front().getResAttrsAttr()[1]).empty()); + + ASSERT_TRUE(succeeded(runQCOToQCConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + std::size_t qcCalls = 0; + moduleOp->walk([&](qc::CallOp call) { + ++qcCalls; + ASSERT_TRUE(call.getArgAttrsAttr()); + EXPECT_EQ(cast(call.getArgAttrsAttr()[0]) + .getAs("tag") + .getValue(), + "unitary-input"); + EXPECT_EQ(call->getAttrOfType("tag").getValue(), + "unitary-call"); + }); + EXPECT_EQ(qcCalls, 1U); + auto genericCall = *mlir::mqt::getEntryPoint(*moduleOp) + .getBody() + .getOps() + .begin(); + EXPECT_TRUE(genericCall.getNoInline()); + EXPECT_EQ(genericCall->getAttrOfType("tag").getValue(), + "generic-call"); + ASSERT_EQ(genericCall.getResAttrsAttr().size(), 1U); + EXPECT_EQ(cast(genericCall.getResAttrsAttr()[0]) + .getAs("tag") + .getValue(), + "ordinary-result"); + for (auto function : moduleOp->getOps()) { + EXPECT_TRUE(isa(function.getArgument(0).getType())); + EXPECT_EQ(function.getNumResults(), function.getName() == "flip" ? 0U : 1U); + } +} + +TEST_F(QCToQCORegressionTest, ConvertsRegisterBackedUnitaryCalls) { + constexpr llvm::StringLiteral source = R"mlir( +module { + func.func private @rotate(%theta: f64, %q: !qc.qubit) + attributes {mqt.unitary} { + qc.rx(%theta) %q : !qc.qubit + return + } + func.func @main() attributes {mqt.entry_point} { + %reg = memref.alloc() : memref<1x!qc.qubit> + %c0 = arith.constant 0 : index + %theta = arith.constant 5.000000e-01 : f64 + %q = memref.load %reg[%c0] : memref<1x!qc.qubit> + qc.call @rotate(%theta, %q) : f64, !qc.qubit + %two = arith.constant 2.000000e+00 : f64 + qc.pow(%two) (%arg0 = %q) { + qc.call @rotate(%theta, %arg0) : f64, !qc.qubit + qc.yield + } : !qc.qubit + memref.dealloc %reg : memref<1x!qc.qubit> + return + } +} +)mlir"; + + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCToQCOConversion(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + EXPECT_TRUE(succeeded(qco::verifyLinearity(*moduleOp))); + size_t calls = 0; + size_t extracts = 0; + size_t inserts = 0; + moduleOp->walk([&](Operation* operation) { + calls += isa(operation); + extracts += isa(operation); + inserts += isa(operation); + }); + EXPECT_EQ(calls, 2); + EXPECT_EQ(extracts, 2); + EXPECT_EQ(inserts, 2); +} + +TEST_F(QCToQCORegressionTest, PreflightRejectsAliasedAndDuplicateQubitResults) { + constexpr auto sources = std::to_array({ + R"mlir( +module { + func.func private @borrowed(%q: !qc.qubit) -> !qc.qubit { + return %q : !qc.qubit + } + func.func @main() attributes {mqt.entry_point} { + return + } +} +)mlir", + R"mlir( +module { + func.func private @duplicate() -> (!qc.qubit, !qc.qubit) { + %q = qc.alloc : !qc.qubit + return %q, %q : !qc.qubit, !qc.qubit + } + func.func @main() attributes {mqt.entry_point} { + return + } +} )mlir", + }); + constexpr std::array diagnostics{ + "cannot return a borrowed qubit argument explicitly", + "cannot return the same qubit more than once"}; + + for (auto [source, expected] : llvm::zip_equal(sources, diagnostics)) { + auto moduleOp = parseSourceString(source, &context); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + OwningOpRef original = cast(moduleOp->clone()); + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(&context, [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= StringRef(diagnostic.str()).contains(expected); + return success(); + }); + EXPECT_TRUE(failed(runQCToQCOConversion(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); + EXPECT_TRUE(OperationEquivalence::isEquivalentTo( + moduleOp->getOperation(), original->getOperation(), + OperationEquivalence::Flags::None)); + } +} + +TEST_F(QCToQCORegressionTest, + PreflightRejectsUnsupportedQuantumRegisterBlockArguments) { + constexpr auto sources = std::to_array({ R"mlir( module { func.func @main(%reg: memref<1x!qc.qubit>) @@ -1513,11 +1729,12 @@ TEST_F(QCToQCORegressionTest, RejectsSameDynamicRegisterIndexWithinOneOperation) { constexpr llvm::StringLiteral source = R"mlir( module { + func.func private @touch(!qc.qubit, !qc.qubit) func.func @main(%i: index) attributes {mqt.entry_point} { %reg = memref.alloc() : memref<2x!qc.qubit> %q0 = memref.load %reg[%i] : memref<2x!qc.qubit> %q1 = memref.load %reg[%i] : memref<2x!qc.qubit> - qc.swap %q0, %q1 : !qc.qubit, !qc.qubit + func.call @touch(%q0, %q1) : (!qc.qubit, !qc.qubit) -> () memref.dealloc %reg : memref<2x!qc.qubit> return } @@ -1613,8 +1830,7 @@ TEST_P(QCToQCOTest, ProgramConversion) { } } -/// \name QCToQCO/QubitManagement/StaticOp.cpp -/// @{ +// QCToQCO/QubitManagement/StaticOp.cpp INSTANTIATE_TEST_SUITE_P( QCStaticOpTest, QCToQCOTest, testing::Values( @@ -1638,10 +1854,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"AllocDeallocPair", MQT_NAMED_BUILDER(qc::allocDeallocPair), MQT_NAMED_BUILDER(qco::emptyQCO)})); -/// @} -/// \name QCToQCO/Modifiers/PowOp.cpp -/// @{ +// QCToQCO/Modifiers/PowOp.cpp INSTANTIATE_TEST_SUITE_P( QCPowOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"CtrlPowSx", @@ -1649,10 +1863,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qco::ctrlPowSx)}, QCToQCOTestCase{"PowTwo", MQT_NAMED_BUILDER(qc::powTwo), MQT_NAMED_BUILDER(qco::powTwo)})); -/// @} -/// \name QCToQCO/Modifiers/CtrlOp.cpp -/// @{ +// QCToQCO/Modifiers/CtrlOp.cpp INSTANTIATE_TEST_SUITE_P( QCCtrlOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"CtrlTwo", MQT_NAMED_BUILDER(qc::ctrlTwo), @@ -1663,10 +1875,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"CtrlInvTwo", MQT_NAMED_BUILDER(qc::ctrlInvTwo), MQT_NAMED_BUILDER(qco::ctrlInvTwo)})); -/// @} -/// \name QCToQCO/Modifiers/InvOp.cpp -/// @{ +// QCToQCO/Modifiers/InvOp.cpp INSTANTIATE_TEST_SUITE_P( QCInvOpTest, QCToQCOTest, testing::Values( @@ -1678,10 +1888,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qco::inverseMultipleControlledIswap)}, QCToQCOTestCase{"InvTwo", MQT_NAMED_BUILDER(qc::invTwo), MQT_NAMED_BUILDER(qco::invTwo)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/BarrierOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/BarrierOp.cpp INSTANTIATE_TEST_SUITE_P( QCBarrierOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"Barrier", MQT_NAMED_BUILDER(qc::barrier), @@ -1693,10 +1901,8 @@ INSTANTIATE_TEST_SUITE_P( "BarrierMultipleQubits", MQT_NAMED_BUILDER(qc::barrierMultipleQubits), MQT_NAMED_BUILDER(qco::barrierMultipleQubits)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/DcxOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/DcxOp.cpp INSTANTIATE_TEST_SUITE_P( QCDCXOpTest, QCToQCOTest, testing::Values( @@ -1708,10 +1914,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledDCX", MQT_NAMED_BUILDER(qc::multipleControlledDcx), MQT_NAMED_BUILDER(qco::multipleControlledDcx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/EcrOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/EcrOp.cpp INSTANTIATE_TEST_SUITE_P( QCECROpTest, QCToQCOTest, testing::Values( @@ -1723,18 +1927,14 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledECR", MQT_NAMED_BUILDER(qc::multipleControlledEcr), MQT_NAMED_BUILDER(qco::multipleControlledEcr)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/GphaseOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/GphaseOp.cpp INSTANTIATE_TEST_SUITE_P(QCGPhaseOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{ "GlobalPhase", MQT_NAMED_BUILDER(qc::globalPhase), MQT_NAMED_BUILDER(qco::globalPhase)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/HOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/HOp.cpp INSTANTIATE_TEST_SUITE_P( QCHOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"H", MQT_NAMED_BUILDER(qc::h), @@ -1749,18 +1949,14 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"HWithoutRegister", MQT_NAMED_BUILDER(qc::hWithoutRegister), MQT_NAMED_BUILDER(qco::hWithoutRegister)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/IdOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/IdOp.cpp INSTANTIATE_TEST_SUITE_P(QCIDOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{ "Identity", MQT_NAMED_BUILDER(qc::identity), MQT_NAMED_BUILDER(qco::alloc1QubitRegister)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/IswapOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/IswapOp.cpp INSTANTIATE_TEST_SUITE_P( QCiSWAPOpTest, QCToQCOTest, testing::Values( @@ -1772,10 +1968,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControllediSWAP", MQT_NAMED_BUILDER(qc::multipleControlledIswap), MQT_NAMED_BUILDER(qco::multipleControlledIswap)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/POp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/POp.cpp INSTANTIATE_TEST_SUITE_P( QCPOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"P", MQT_NAMED_BUILDER(qc::p), @@ -1787,10 +1981,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledP", MQT_NAMED_BUILDER(qc::multipleControlledP), MQT_NAMED_BUILDER(qco::multipleControlledP)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RCCXOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RCCXOp.cpp INSTANTIATE_TEST_SUITE_P( QCRCCXOpTest, QCToQCOTest, testing::Values( @@ -1802,10 +1994,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledRCCX", MQT_NAMED_BUILDER(qc::multipleControlledRccx), MQT_NAMED_BUILDER(qco::multipleControlledRccx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/ROp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/ROp.cpp INSTANTIATE_TEST_SUITE_P( QCROpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"R", MQT_NAMED_BUILDER(qc::r), @@ -1817,10 +2007,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledR", MQT_NAMED_BUILDER(qc::multipleControlledR), MQT_NAMED_BUILDER(qco::multipleControlledR)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RxOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRXOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"RX", MQT_NAMED_BUILDER(qc::rx), @@ -1832,10 +2020,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRX", MQT_NAMED_BUILDER(qc::multipleControlledRx), MQT_NAMED_BUILDER(qco::multipleControlledRx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RxxOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RxxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRXXOpTest, QCToQCOTest, testing::Values( @@ -1847,10 +2033,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledRXX", MQT_NAMED_BUILDER(qc::multipleControlledRxx), MQT_NAMED_BUILDER(qco::multipleControlledRxx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RyOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RyOp.cpp INSTANTIATE_TEST_SUITE_P( QCRYOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"RY", MQT_NAMED_BUILDER(qc::ry), @@ -1862,10 +2046,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRY", MQT_NAMED_BUILDER(qc::multipleControlledRy), MQT_NAMED_BUILDER(qco::multipleControlledRy)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RyyOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RyyOp.cpp INSTANTIATE_TEST_SUITE_P( QCRYYOpTest, QCToQCOTest, testing::Values( @@ -1877,10 +2059,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledRYY", MQT_NAMED_BUILDER(qc::multipleControlledRyy), MQT_NAMED_BUILDER(qco::multipleControlledRyy)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RzOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RzOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"RZ", MQT_NAMED_BUILDER(qc::rz), @@ -1892,10 +2072,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledRZ", MQT_NAMED_BUILDER(qc::multipleControlledRz), MQT_NAMED_BUILDER(qco::multipleControlledRz)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RzxOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RzxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZXOpTest, QCToQCOTest, testing::Values( @@ -1907,10 +2085,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledRZX", MQT_NAMED_BUILDER(qc::multipleControlledRzx), MQT_NAMED_BUILDER(qco::multipleControlledRzx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/RzzOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/RzzOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZZOpTest, QCToQCOTest, testing::Values( @@ -1922,10 +2098,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledRZZ", MQT_NAMED_BUILDER(qc::multipleControlledRzz), MQT_NAMED_BUILDER(qco::multipleControlledRzz)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/SOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/SOp.cpp INSTANTIATE_TEST_SUITE_P( QCSOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"S", MQT_NAMED_BUILDER(qc::s), @@ -1937,10 +2111,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledS", MQT_NAMED_BUILDER(qc::multipleControlledS), MQT_NAMED_BUILDER(qco::multipleControlledS)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/SdgOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/SdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCSdgOpTest, QCToQCOTest, testing::Values( @@ -1952,10 +2124,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledSdg", MQT_NAMED_BUILDER(qc::multipleControlledSdg), MQT_NAMED_BUILDER(qco::multipleControlledSdg)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/SwapOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/SwapOp.cpp INSTANTIATE_TEST_SUITE_P( QCSWAPOpTest, QCToQCOTest, testing::Values( @@ -1967,10 +2137,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledSWAP", MQT_NAMED_BUILDER(qc::multipleControlledSwap), MQT_NAMED_BUILDER(qco::multipleControlledSwap)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/SxOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/SxOp.cpp INSTANTIATE_TEST_SUITE_P( QCSXOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"SX", MQT_NAMED_BUILDER(qc::sx), @@ -1982,10 +2150,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledSX", MQT_NAMED_BUILDER(qc::multipleControlledSx), MQT_NAMED_BUILDER(qco::multipleControlledSx)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/SxdgOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/SxdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCSXdgOpTest, QCToQCOTest, testing::Values( @@ -1997,10 +2163,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledSXdg", MQT_NAMED_BUILDER(qc::multipleControlledSxdg), MQT_NAMED_BUILDER(qco::multipleControlledSxdg)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/TOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/TOp.cpp INSTANTIATE_TEST_SUITE_P( QCTOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"T", MQT_NAMED_BUILDER(qc::t_), @@ -2012,10 +2176,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledT", MQT_NAMED_BUILDER(qc::multipleControlledT), MQT_NAMED_BUILDER(qco::multipleControlledT)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/TdgOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/TdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCTdgOpTest, QCToQCOTest, testing::Values( @@ -2027,10 +2189,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledTdg", MQT_NAMED_BUILDER(qc::multipleControlledTdg), MQT_NAMED_BUILDER(qco::multipleControlledTdg)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/U2Op.cpp -/// @{ +// QCToQCO/Operations/StandardGates/U2Op.cpp INSTANTIATE_TEST_SUITE_P( QCU2OpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"U2", MQT_NAMED_BUILDER(qc::u2), @@ -2042,10 +2202,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledU2", MQT_NAMED_BUILDER(qc::multipleControlledU2), MQT_NAMED_BUILDER(qco::multipleControlledU2)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/UOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/UOp.cpp INSTANTIATE_TEST_SUITE_P( QCUOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"U", MQT_NAMED_BUILDER(qc::u), @@ -2057,10 +2215,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledU", MQT_NAMED_BUILDER(qc::multipleControlledU), MQT_NAMED_BUILDER(qco::multipleControlledU)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/XOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/XOp.cpp INSTANTIATE_TEST_SUITE_P( QCXOpTest, QCToQCOTest, testing::Values( @@ -2076,10 +2232,7 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qc::repeatedControlledX), MQT_NAMED_BUILDER(qco::repeatedControlledX)})); -/// @} - -/// \name QCToQCO/Operations/StandardGates/XxMinusYyOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/XxMinusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCXXMinusYYOpTest, QCToQCOTest, testing::Values( @@ -2091,10 +2244,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledXXMinusYY", MQT_NAMED_BUILDER(qc::multipleControlledXxMinusYY), MQT_NAMED_BUILDER(qco::multipleControlledXxMinusYY)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/XxPlusYyOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/XxPlusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCXXPlusYYOpTest, QCToQCOTest, testing::Values( @@ -2106,10 +2257,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MultipleControlledXXPlusYY", MQT_NAMED_BUILDER(qc::multipleControlledXxPlusYY), MQT_NAMED_BUILDER(qco::multipleControlledXxPlusYY)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/YOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/YOp.cpp INSTANTIATE_TEST_SUITE_P( QCYOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"Y", MQT_NAMED_BUILDER(qc::y), @@ -2121,10 +2270,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledY", MQT_NAMED_BUILDER(qc::multipleControlledY), MQT_NAMED_BUILDER(qco::multipleControlledY)})); -/// @} -/// \name QCToQCO/Operations/StandardGates/ZOp.cpp -/// @{ +// QCToQCO/Operations/StandardGates/ZOp.cpp INSTANTIATE_TEST_SUITE_P( QCZOpTest, QCToQCOTest, testing::Values(QCToQCOTestCase{"Z", MQT_NAMED_BUILDER(qc::z), @@ -2136,10 +2283,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleControlledZ", MQT_NAMED_BUILDER(qc::multipleControlledZ), MQT_NAMED_BUILDER(qco::multipleControlledZ)})); -/// @} -/// \name QCToQCO/Operations/MeasureOp.cpp -/// @{ +// QCToQCO/Operations/MeasureOp.cpp INSTANTIATE_TEST_SUITE_P( QCMeasureOpTest, QCToQCOTest, testing::Values( @@ -2167,10 +2312,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"MeasurementWithoutRegisters", MQT_NAMED_BUILDER(qc::measurementWithoutRegisters), MQT_NAMED_BUILDER(qco::measurementWithoutRegisters)})); -/// @} -/// \name QCToQCO/Operations/ResetOp.cpp -/// @{ +// QCToQCO/Operations/ResetOp.cpp INSTANTIATE_TEST_SUITE_P( QCResetOpTest, QCToQCOTest, testing::Values( @@ -2184,10 +2327,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"RepeatedResetAfterSingleOp", MQT_NAMED_BUILDER(qc::repeatedResetAfterSingleOp), MQT_NAMED_BUILDER(qco::resetQubitAfterSingleOp)})); -/// @} -/// \name QCToQCO/Operations/IfOp.cpp -/// @{ +// QCToQCO/Operations/IfOp.cpp INSTANTIATE_TEST_SUITE_P( SCFIfOpTest, QCToQCOTest, testing::Values( @@ -2209,10 +2350,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"NestedIfOpForLoop", MQT_NAMED_BUILDER(qc::nestedIfOpForLoop), MQT_NAMED_BUILDER(qco::nestedIfOpForLoop), true})); -/// @} -/// \name QCToQCO/Operations/IndexSwitchOp.cpp -/// @{ +// QCToQCO/Operations/IndexSwitchOp.cpp INSTANTIATE_TEST_SUITE_P( QCOIndexSwitchOpTest, QCToQCOTest, testing::Values( @@ -2223,10 +2362,8 @@ INSTANTIATE_TEST_SUITE_P( "IndexSwitchMultiCase", MQT_NAMED_BUILDER(qc::indexSwitchMultiCase), MQT_NAMED_BUILDER(qco::indexSwitchMultiCaseCompleteTensorState), true})); -/// @} -/// \name QCToQCO/Operations/WhileOp.cpp -/// @{ +// QCToQCO/Operations/WhileOp.cpp INSTANTIATE_TEST_SUITE_P( SCFWhileOpTest, QCToQCOTest, testing::Values( @@ -2235,10 +2372,8 @@ INSTANTIATE_TEST_SUITE_P( QCToQCOTestCase{"SimpleDoWhile", MQT_NAMED_BUILDER(qc::simpleDoWhileReset), MQT_NAMED_BUILDER(qco::simpleDoWhileReset)})); -/// @} -/// \name QCToQCO/Operations/ForOp.cpp -/// @{ +// QCToQCO/Operations/ForOp.cpp INSTANTIATE_TEST_SUITE_P( SCFForOpTest, QCToQCOTest, testing::Values( @@ -2265,4 +2400,3 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(qc::nestedForLoopCtrlOpWithExtractedQubit), MQT_NAMED_BUILDER(qco::nestedForLoopCtrlOpWithExtractedQubit), true})); -/// @} diff --git a/mlir/unittests/Dialect/MQT/IR/test_mqt_ir.cpp b/mlir/unittests/Dialect/MQT/IR/test_mqt_ir.cpp index 63edc58e5c..9af028efe4 100644 --- a/mlir/unittests/Dialect/MQT/IR/test_mqt_ir.cpp +++ b/mlir/unittests/Dialect/MQT/IR/test_mqt_ir.cpp @@ -8,10 +8,8 @@ * Licensed under the MIT License */ -/** - * @file test_mqt_ir.cpp - * @brief Unit tests for the MQT metadata dialect. - */ +/// @file test_mqt_ir.cpp +/// Unit tests for the MQT metadata dialect. #include "mlir/Dialect/CBit/IR/CBitDialect.h" #include "mlir/Dialect/MQT/IR/MQTAttributes.h" @@ -29,6 +27,7 @@ #include #include #include +#include #include #include #include @@ -102,6 +101,36 @@ TEST_F(MQTIRTest, AcceptsProgramInputAndRegisterNames) { )mlir")); } +TEST_F(MQTIRTest, AcceptsSourceFunctionName) { + EXPECT_TRUE(parse(R"mlir( + module { + func.func private @unique() attributes {mqt.source_name = "source"} + } + )mlir")); +} + +TEST_F(MQTIRTest, RejectsInvalidSourceFunctionNames) { + EXPECT_FALSE(parse(R"mlir( + module { + func.func private @empty() attributes {mqt.source_name = ""} + } + )mlir")); + EXPECT_FALSE(parse(R"mlir( + module { + func.func private @null() attributes {mqt.source_name = "a\00b"} + } + )mlir")); + EXPECT_FALSE(parse(R"mlir( + module { + func.func @main() { + %c0 = "arith.constant"() {mqt.source_name = "source", value = 0 : i64} + : () -> i64 + return + } + } + )mlir")); +} + TEST_F(MQTIRTest, RoundTripsTypedCompilationTarget) { const auto compilationTarget = dyn_cast_if_present( @@ -303,6 +332,164 @@ TEST_F(MQTIRTest, RejectsInvalidEntryPoints) { )mlir")); } +TEST_F(MQTIRTest, RejectsMutuallyRecursiveUnitaryFunctions) { + for (StringRef source : { + R"mlir( + func.func private @first(%q: !qc.qubit) attributes {mqt.unitary} { + qc.call @second(%q) : !qc.qubit + return + } + func.func private @second(%q: !qc.qubit) attributes {mqt.unitary} { + qc.call @first(%q) : !qc.qubit + return + } + )mlir", + R"mlir( + func.func private @first(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.call @second(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + func.func private @second(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.call @first(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + )mlir"}) { + SCOPED_TRACE(source.str()); + bool sawRecursion = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawRecursion |= StringRef(diagnostic.str()) + .contains("unitary function must not be recursive"); + return success(); + }); + EXPECT_FALSE(parse(source)); + EXPECT_TRUE(sawRecursion); + } +} + +TEST_F(MQTIRTest, RejectsEmptyUnitaryBodies) { + for (StringRef source : { + R"mlir( + func.func private @empty(!qc.qubit) attributes {mqt.unitary} { + ^bb0(%q: !qc.qubit): + } + )mlir", + R"mlir( + func.func private @empty(!qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + ^bb0(%q: !qco.qubit): + } + )mlir"}) { + SCOPED_TRACE(source.str()); + bool sawEmptyBody = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawEmptyBody |= StringRef(diagnostic.str()) + .contains("unitary function body must not be empty"); + return success(); + }); + EXPECT_FALSE(parse(source)); + EXPECT_TRUE(sawEmptyBody); + } +} + +TEST_F(MQTIRTest, RejectsCyclicUnitaryQubitFlow) { + bool sawCycle = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawCycle |= StringRef(diagnostic.str()) + .contains("unitary QCO result has cyclic qubit flow"); + return success(); + }); + EXPECT_FALSE(parse(R"mlir( + func.func private @cyclic(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %a = qco.h %b : !qco.qubit -> !qco.qubit + %b = qco.h %a : !qco.qubit -> !qco.qubit + return %b : !qco.qubit + } + )mlir")); + EXPECT_TRUE(sawCycle); +} + +TEST_F(MQTIRTest, RejectsMalformedUnitaryBodyOperations) { + for (StringRef source : { + R"mlir( + func.func private @malformed(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = "qco.h"() : () -> !qco.qubit + return %out : !qco.qubit + } + )mlir", + R"mlir( + func.func private @malformed(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.inv (%arg = %q) { + %h = "qco.h"() : () -> !qco.qubit + qco.yield %h : !qco.qubit + } : {!qco.qubit} -> {!qco.qubit} + return %out : !qco.qubit + } + )mlir", + R"mlir( + func.func private @malformed(%q: !qc.qubit) + attributes {mqt.unitary} { + %value = "memref.load"() : () -> f64 + return + } + )mlir", + R"mlir( + func.func private @malformed(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %value = "memref.load"() : () -> f64 + return %q : !qco.qubit + } + )mlir"}) { + SCOPED_TRACE(source.str()); + bool sawOperandError = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawOperandError |= StringRef(diagnostic.str()).contains("operand"); + return success(); + }); + EXPECT_FALSE(parse(source)); + EXPECT_TRUE(sawOperandError); + } +} + +TEST_F(MQTIRTest, RejectsMalformedCallsInUnitaryCallees) { + for (StringRef source : { + R"mlir( + func.func private @first(%q: !qc.qubit) attributes {mqt.unitary} { + qc.call @second(%q) : !qc.qubit + return + } + func.func private @second(%q: !qc.qubit) attributes {mqt.unitary} { + "qc.call"(%q) : (!qc.qubit) -> () + return + } + )mlir", + R"mlir( + func.func private @first(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.call @second(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + func.func private @second(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = "qco.call"(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + )mlir"}) { + SCOPED_TRACE(source.str()); + bool sawMissingCallee = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawMissingCallee |= StringRef(diagnostic.str()).contains("callee"); + return success(); + }); + EXPECT_FALSE(parse(source)); + EXPECT_TRUE(sawMissingCallee); + } +} + TEST_F(MQTIRTest, RejectsInvalidInputNames) { EXPECT_FALSE(parse(R"mlir( module { diff --git a/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp b/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp index f3561243bd..508fd87646 100644 --- a/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp +++ b/mlir/unittests/Dialect/QC/IR/test_qc_ir.cpp @@ -133,17 +133,17 @@ TEST_P(QCTest, ProgramEquivalence) { } TEST_F(QCTest, QubitIsVectorElement) { - auto module = parseSourceString(R"mlir( + auto moduleOp = parseSourceString(R"mlir( module { func.func @f(%arg: vector<2x!qc.qubit>) { return } } )mlir", - context.get()); - ASSERT_TRUE(module); + context.get()); + ASSERT_TRUE(moduleOp); - auto function = *module->getOps().begin(); + auto function = *moduleOp->getOps().begin(); const auto vectorType = dyn_cast(function.getArgument(0).getType()); ASSERT_TRUE(vectorType); @@ -151,7 +151,7 @@ TEST_F(QCTest, QubitIsVectorElement) { } TEST_F(QCTest, CleanupHoistsAndCoalescesStaticQubits) { - auto module = parseSourceString(R"mlir( + auto moduleOp = parseSourceString(R"mlir( module { func.func @main(%condition: i1) { scf.if %condition { @@ -167,14 +167,14 @@ TEST_F(QCTest, CleanupHoistsAndCoalescesStaticQubits) { } } )mlir", - context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(runQCCleanupPipeline(*module))); + context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(runQCCleanupPipeline(*moduleOp))); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); size_t staticOps = 0; - module->walk([&](StaticOp op) { + moduleOp->walk([&](StaticOp op) { ++staticOps; EXPECT_EQ(op->getBlock(), &main.getBody().front()); }); @@ -182,7 +182,7 @@ TEST_F(QCTest, CleanupHoistsAndCoalescesStaticQubits) { } TEST_F(QCTest, CleanupDoesNotHoistStaticQubitsAcrossIsolationBoundaries) { - auto module = parseSourceString(R"mlir( + auto moduleOp = parseSourceString(R"mlir( module { func.func @main() { "builtin.module"() ({ @@ -193,14 +193,14 @@ TEST_F(QCTest, CleanupDoesNotHoistStaticQubitsAcrossIsolationBoundaries) { } } )mlir", - context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCCleanupPipeline(*module))); - ASSERT_TRUE(succeeded(verify(*module))); + context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCCleanupPipeline(*moduleOp))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); StaticOp staticOp; - module->walk([&](StaticOp op) { staticOp = op; }); + moduleOp->walk([&](StaticOp op) { staticOp = op; }); ASSERT_TRUE(staticOp); EXPECT_TRUE(isa(staticOp->getParentOp())); } @@ -221,6 +221,30 @@ TEST_F(QCTest, BuilderRejectsMixedStaticAndDynamicQubitAllocationModes) { mixedDynamicRegisterThenStaticQubit(builder); }, "Cannot mix dynamic and static qubit allocation modes"); + + EXPECT_DEATH( + { + QCProgramBuilder builder(context.get()); + builder.initialize(); + builder.allocQubit(); + builder.createFunction("static_helper", {}, [&](ValueRange) { + builder.staticQubit(0); + return SmallVector{}; + }); + }, + "Cannot mix dynamic and static qubit allocation modes"); + + EXPECT_DEATH( + { + QCProgramBuilder builder(context.get()); + builder.initialize(); + builder.createFunction("dynamic_helper", {}, [&](ValueRange) { + builder.allocQubit(); + return SmallVector{}; + }); + builder.staticQubit(0); + }, + "Cannot mix dynamic and static qubit allocation modes"); } TEST_F(QCTest, BuilderRejectsOutOfBoundsClassicalRegisterIndices) { @@ -369,6 +393,238 @@ TEST_F(QCTest, BuilderCanAllocateQubitRegisterStorageWithoutEagerLoads) { EXPECT_EQ(qubitLoads, 0U); } +TEST_F(QCTest, BuilderCreatesGenericAndUnitaryFunctions) { + QCProgramBuilder builder(context.get()); + builder.initialize(); + + auto generic = builder.createFunction( + "identity", TypeRange{builder.getI1Type()}, + [](ValueRange arguments) { return SmallVector{arguments[0]}; }); + auto unitary = builder.createUnitaryFunction( + "flip", TypeRange{QubitType::get(context.get())}, + [&](ValueRange arguments) { builder.x(arguments[0]); }); + + auto bit = builder.boolConstant(true); + auto genericResults = builder.call(generic, bit); + auto qubit = builder.allocQubit(); + EXPECT_TRUE(builder.call(unitary, qubit).empty()); + builder.inv(qubit, [&](Value argument) { builder.call(unitary, argument); }); + builder.retype(ValueRange(genericResults).getTypes()); + auto moduleOp = builder.finalize(genericResults); + + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + EXPECT_TRUE(generic.isPrivate()); + EXPECT_FALSE(mlir::mqt::isUnitaryFunction(generic)); + EXPECT_TRUE(mlir::mqt::isUnitaryFunction(unitary)); + EXPECT_EQ(generic.getNumResults(), 1U); + EXPECT_EQ(unitary.getNumResults(), 0U); + + auto mainFunc = mlir::mqt::getEntryPoint(*moduleOp); + ASSERT_TRUE(mainFunc); + EXPECT_EQ(llvm::range_size(mainFunc.getBody().getOps()), 1U); + EXPECT_EQ(llvm::range_size(mainFunc.getBody().getOps()), 1U); + auto directCall = *mainFunc.getBody().getOps().begin(); + ASSERT_EQ(directCall.getQubits().size(), 1U); + EXPECT_EQ(directCall.getQubits().front(), qubit); + auto inverse = *mainFunc.getBody().getOps().begin(); + auto nestedCall = *inverse.getRegion().getOps().begin(); + EXPECT_TRUE(isa(nestedCall.getOperation())); +} + +TEST_F(QCTest, BuilderFinalizesRenamedEntryPoint) { + QCProgramBuilder builder(context.get()); + builder.initialize(); + auto entry = cast(builder.getInsertionBlock()->getParentOp()); + entry.setName("entry"); + + auto moduleOp = builder.finalize(); + + ASSERT_TRUE(moduleOp); + EXPECT_EQ(mlir::mqt::getEntryPoint(*moduleOp).getName(), "entry"); +} + +TEST_F(QCTest, BuilderCreatesFunctionLocalStaticQubits) { + QCProgramBuilder builder(context.get()); + builder.initialize(); + auto helper = builder.createFunction("helper", {}, [&](ValueRange) { + builder.x(builder.staticQubit(0)); + return SmallVector{}; + }); + auto mainQubit = builder.staticQubit(0); + builder.x(mainQubit); + auto moduleOp = builder.finalize(); + + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + auto main = mlir::mqt::getEntryPoint(*moduleOp); + ASSERT_TRUE(main); + EXPECT_EQ(llvm::range_size(helper.getOps()), 1U); + EXPECT_EQ(llvm::range_size(main.getOps()), 1U); +} + +TEST_F(QCTest, BuilderRejectsInvalidCalls) { + EXPECT_DEATH( + { + QCProgramBuilder builder(context.get()); + builder.initialize(); + auto function = builder.createFunction( + "identity", TypeRange{builder.getI1Type()}, + [](ValueRange arguments) { + return SmallVector{arguments.front()}; + }); + builder.call(function, builder.intConstant(0)); + }, + "Call operands must match a function in the current module"); + + EXPECT_DEATH( + { + QCProgramBuilder first(context.get()); + first.initialize(); + QCProgramBuilder second(context.get()); + second.initialize(); + auto function = second.createFunction( + "identity", TypeRange{second.getI1Type()}, + [](ValueRange arguments) { + return SmallVector{arguments.front()}; + }); + first.call(function, first.boolConstant(true)); + }, + "Call operands must match a function in the current module"); +} + +TEST_F(QCTest, UnitaryFunctionMarkerRejectsNonUnitaryBody) { + QCProgramBuilder builder(context.get()); + builder.initialize(); + auto function = builder.createFunction( + "measure", TypeRange{QubitType::get(context.get())}, + [&](ValueRange arguments) { + builder.measure(arguments[0]); + return SmallVector{}; + }); + mlir::mqt::setUnitaryFunction(function); + auto moduleOp = builder.finalize(); + + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= + StringRef(diagnostic.str()) + .contains( + "unitary QC function body contains a non-unitary operation"); + return success(); + }); + EXPECT_TRUE(failed(verify(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + +TEST_F(QCTest, UnitaryFunctionMarkerRequiresFunctionReturn) { + QCProgramBuilder builder(context.get()); + builder.initialize(); + auto function = builder.createUnitaryFunction( + "flip", TypeRange{QubitType::get(context.get())}, + [&](ValueRange arguments) { builder.x(arguments[0]); }); + auto moduleOp = builder.finalize(); + + auto returnOp = cast(function.getBody().front().back()); + OpBuilder rewriter(returnOp); + YieldOp::create(rewriter, returnOp.getLoc()); + returnOp.erase(); + + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= + StringRef(diagnostic.str()) + .contains("unitary QC function must end in an empty func.return"); + return success(); + }); + EXPECT_TRUE(failed(verify(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + +TEST_F(QCTest, UnitaryVerifierRejectsInvalidFunctionAndCallContracts) { + DialectRegistry registry; + registry.insert(); + context->appendDialectRegistry(registry); + context->getOrLoadDialect(); + + constexpr std::array invalidPrograms{ + R"mlir(module { + func.func private @bad(%q: !qc.qubit) + attributes {mqt.unitary = true} { return } + })mlir", + R"mlir(module { + func.func @bad(%q: !qc.qubit) attributes {mqt.unitary} { return } + })mlir", + R"mlir(module { + func.func private @bad() attributes {mqt.unitary} { return } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qc.qubit, %theta: f64) + attributes {mqt.unitary} { return } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qc.qubit) -> i1 + attributes {mqt.unitary} { + %value = arith.constant true + return %value : i1 + } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qc.qubit) attributes {mqt.unitary} { + qc.call @bad(%q) : !qc.qubit + return + } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qc.qubit) attributes {mqt.unitary} { + qc.call @missing(%q) : !qc.qubit + return + } + })mlir", + R"mlir(module { + func.func private @plain(%q: !qc.qubit) { return } + func.func @main(%q: !qc.qubit) attributes {mqt.entry_point} { + qc.call @plain(%q) : !qc.qubit + return + } + })mlir", + R"mlir(module { + func.func private @flip(%q: !qc.qubit) attributes {mqt.unitary} { + qc.x %q : !qc.qubit + return + } + func.func @main(%theta: f64, %q: !qc.qubit) + attributes {mqt.entry_point} { + qc.call @flip(%theta, %q) : f64, !qc.qubit + return + } + })mlir", + }; + + ParserConfig config(context.get(), false); + for (const auto source : invalidPrograms) { + auto moduleOp = parseSourceString(source, config); + ASSERT_TRUE(moduleOp); + EXPECT_TRUE(failed(verify(*moduleOp))); + } + + auto resultModule = parseSourceString(R"mlir(module { + func.func private @bad(%q: !qc.qubit) -> i1 attributes {mqt.unitary} + func.func @main(%q: !qc.qubit) attributes {mqt.entry_point} { + qc.call @bad(%q) : !qc.qubit + return + } + })mlir", + config); + ASSERT_TRUE(resultModule); + auto call = *mlir::mqt::getEntryPoint(*resultModule) + .getBody() + .getOps() + .begin(); + SymbolTableCollection symbols; + EXPECT_TRUE(failed(call.verifySymbolUses(symbols))); +} + TEST_F(QCTest, DirectSingleQubitPowBuilder) { QCProgramBuilder builder(context.get()); builder.initialize(); @@ -441,16 +697,16 @@ TEST_F(QCTest, DenseUnitaryBuilderVerifiesAndCanonicalizesIdentity) { builder.initialize(); auto qubit = builder.allocQubit(); builder.unitary(ValueRange{qubit}, xMatrix); - auto module = builder.finalize(); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = builder.finalize(); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); - auto function = *module->getOps().begin(); + auto function = *moduleOp->getOps().begin(); auto unitaries = llvm::to_vector(function.getBody().getOps()); ASSERT_EQ(unitaries.size(), 1U); EXPECT_EQ(unitaries.front().getMatrix(), xMatrix); - ASSERT_TRUE(succeeded(runQCCleanupPipeline(*module))); + ASSERT_TRUE(succeeded(runQCCleanupPipeline(*moduleOp))); unitaries = llvm::to_vector(function.getBody().getOps()); ASSERT_EQ(unitaries.size(), 1U); EXPECT_EQ(unitaries.front().getMatrix(), xMatrix); @@ -461,7 +717,7 @@ TEST_F(QCTest, DenseUnitaryBuilderVerifiesAndCanonicalizesIdentity) { "matrix", DenseElementsAttr::get( matrixType, llvm::ArrayRef>(identityValues))); - ASSERT_TRUE(succeeded(runQCCleanupPipeline(*module))); + ASSERT_TRUE(succeeded(runQCCleanupPipeline(*moduleOp))); EXPECT_TRUE(function.getBody().getOps().empty()); } @@ -846,8 +1102,7 @@ TEST_F(QCTest, ModifiersRejectDirectAndNestedQubitCaptures) { } } -/// \name QC/Modifiers/CtrlOp.cpp -/// @{ +// QC/Modifiers/CtrlOp.cpp INSTANTIATE_TEST_SUITE_P( QCCtrlOpTest, QCTest, testing::Values( @@ -869,7 +1124,6 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"ModifierBodyReuseReordered", MQT_NAMED_BUILDER(modifierBodyReuseReordered), MQT_NAMED_BUILDER(modifierBodyReuseReorderedRef)})); -/// @} /// A power modifier with a qubit that its body does not use. static Value powWithUnusedQubit(QCProgramBuilder& b) { @@ -878,8 +1132,7 @@ static Value powWithUnusedQubit(QCProgramBuilder& b) { return measureRegister(b, q.qubits); } -/// \name QC/Modifiers/PowOp.cpp -/// @{ +// QC/Modifiers/PowOp.cpp INSTANTIATE_TEST_SUITE_P( QCPowOpTest, QCTest, testing::Values( @@ -920,7 +1173,6 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(ctrlPowSxRef)}, QCTestCase{"PowWithUnusedQubit", MQT_NAMED_BUILDER(powWithUnusedQubit), MQT_NAMED_BUILDER(twoQubitsOneBarrier)})); -/// @} TEST_F(QCTest, PowExponentIsUnitaryParameter) { auto program = @@ -1002,8 +1254,8 @@ TEST_F(QCTest, NestedPowAcrossBranchCutDoesNotMerge) { EXPECT_EQ(xCount, 0); } -/// pow(-0.5) { h } cannot fold a negative fractional exponent -/// into H (no angle to scale). Verify that PowOp survives. +// pow(-0.5) { h } cannot fold a negative fractional exponent +// into H (no angle to scale). Verify that PowOp survives. TEST_F(QCTest, NegPowHNoFold) { auto program = ::mqt::test::buildMLIRProgram(context.get(), MQT_NAMED_BUILDER(negPowH)); @@ -1017,8 +1269,8 @@ TEST_F(QCTest, NegPowHNoFold) { EXPECT_EQ(powCount, 1) << "PowOp around h must survive the pipeline"; } -/// A multi-unitary pow body (pow(2){x; rxx}) is left untouched by the cleanup -/// pipeline. Verify the pow and both body unitaries survive. +// A multi-unitary pow body (pow(2){x; rxx}) is left untouched by the cleanup +// pipeline. Verify the pow and both body unitaries survive. TEST_F(QCTest, PowTwoSurvives) { auto program = ::mqt::test::buildMLIRProgram(context.get(), MQT_NAMED_BUILDER(powTwo)); @@ -1037,8 +1289,7 @@ TEST_F(QCTest, PowTwoSurvives) { EXPECT_EQ(bodyUnitaries, 2U) << "both body unitaries must be preserved"; } -/// \name QC/Modifiers/InvOp.cpp -/// @{ +// QC/Modifiers/InvOp.cpp INSTANTIATE_TEST_SUITE_P( QCInvOpTest, QCTest, testing::Values(QCTestCase{"EmptyInv", MQT_NAMED_BUILDER(emptyInv), @@ -1053,10 +1304,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(singleControlledRxx)}, QCTestCase{"InverseT", MQT_NAMED_BUILDER(inverseT), MQT_NAMED_BUILDER(tdg)})); -/// @} -/// \name QC/Operations/MeasureOp.cpp -/// @{ +// QC/Operations/MeasureOp.cpp INSTANTIATE_TEST_SUITE_P( QCMeasureOpTest, QCTest, testing::Values( @@ -1073,10 +1322,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleClassicalRegistersAndMeasurements", MQT_NAMED_BUILDER(multipleClassicalRegistersAndMeasurements), MQT_NAMED_BUILDER(multipleClassicalRegistersAndMeasurements)})); -/// @} -/// \name QC/Operations/ResetOp.cpp -/// @{ +// QC/Operations/ResetOp.cpp INSTANTIATE_TEST_SUITE_P( QCResetOpTest, QCTest, testing::Values(QCTestCase{"ResetQubitWithoutOp", @@ -1098,10 +1345,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"RepeatedResetAfterSingleOp", MQT_NAMED_BUILDER(repeatedResetAfterSingleOp), MQT_NAMED_BUILDER(repeatedResetAfterSingleOp)})); -/// @} -/// \name QC/Operations/StandardGates/BarrierOp.cpp -/// @{ +// QC/Operations/StandardGates/BarrierOp.cpp INSTANTIATE_TEST_SUITE_P( QCBarrierOpTest, QCTest, testing::Values(QCTestCase{"Barrier", MQT_NAMED_BUILDER(barrier), @@ -1120,10 +1365,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(barrier)}, QCTestCase{"PowBarrier", MQT_NAMED_BUILDER(powBarrier), MQT_NAMED_BUILDER(barrier)})); -/// @} -/// \name QC/Operations/StandardGates/DcxOp.cpp -/// @{ +// QC/Operations/StandardGates/DcxOp.cpp INSTANTIATE_TEST_SUITE_P( QCDCXOpTest, QCTest, testing::Values(QCTestCase{"DCX", MQT_NAMED_BUILDER(dcx), @@ -1145,10 +1388,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledDCX", MQT_NAMED_BUILDER(inverseMultipleControlledDcx), MQT_NAMED_BUILDER(multipleControlledDcx)})); -/// @} -/// \name QC/Operations/StandardGates/EcrOp.cpp -/// @{ +// QC/Operations/StandardGates/EcrOp.cpp INSTANTIATE_TEST_SUITE_P( QCECROpTest, QCTest, testing::Values(QCTestCase{"ECR", MQT_NAMED_BUILDER(ecr), @@ -1174,10 +1415,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(allocQubitRegister)}, QCTestCase{"PowOddECR", MQT_NAMED_BUILDER(powOddEcr), MQT_NAMED_BUILDER(ecr)})); -/// @} -/// \name QC/Operations/StandardGates/GphaseOp.cpp -/// @{ +// QC/Operations/StandardGates/GphaseOp.cpp INSTANTIATE_TEST_SUITE_P( QCGPhaseOpTest, QCTest, testing::Values( @@ -1204,10 +1443,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powGphaseScaledRef)}, QCTestCase{"NegPowGphase", MQT_NAMED_BUILDER(negPowGphase), MQT_NAMED_BUILDER(negPowGphaseRef)})); -/// @} -/// \name QC/Operations/StandardGates/HOp.cpp -/// @{ +// QC/Operations/StandardGates/HOp.cpp INSTANTIATE_TEST_SUITE_P( QCHOpTest, QCTest, testing::Values( @@ -1230,10 +1467,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc1QubitRegister)}, QCTestCase{"PowOddH", MQT_NAMED_BUILDER(powOddH), MQT_NAMED_BUILDER(h)})); -/// @} -/// \name QC/Operations/StandardGates/IdOp.cpp -/// @{ +// QC/Operations/StandardGates/IdOp.cpp INSTANTIATE_TEST_SUITE_P( QCIDOpTest, QCTest, testing::Values( @@ -1258,10 +1493,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(threeQubitsOneIdentity)}, QCTestCase{"PowId", MQT_NAMED_BUILDER(powId), MQT_NAMED_BUILDER(identity)})); -/// @} -/// \name QC/Operations/StandardGates/IswapOp.cpp -/// @{ +// QC/Operations/StandardGates/IswapOp.cpp INSTANTIATE_TEST_SUITE_P( QCiSWAPOpTest, QCTest, testing::Values( @@ -1285,10 +1518,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(inverseMultipleControlledIswap)}, QCTestCase{"PowHalfiSWAP", MQT_NAMED_BUILDER(powHalfIswap), MQT_NAMED_BUILDER(powHalfIswapRef)})); -/// @} -/// \name QC/Operations/StandardGates/POp.cpp -/// @{ +// QC/Operations/StandardGates/POp.cpp INSTANTIATE_TEST_SUITE_P( QCPOpTest, QCTest, testing::Values( @@ -1307,10 +1538,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledP", MQT_NAMED_BUILDER(inverseMultipleControlledP), MQT_NAMED_BUILDER(multipleControlledP)})); -/// @} -/// \name QC/Operations/StandardGates/RCCXOp.cpp -/// @{ +// QC/Operations/StandardGates/RCCXOp.cpp INSTANTIATE_TEST_SUITE_P( QCRCCXOpTest, QCTest, testing::Values(QCTestCase{"RCCX", MQT_NAMED_BUILDER(rccx), @@ -1336,10 +1565,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRCCX", MQT_NAMED_BUILDER(inverseMultipleControlledRccx), MQT_NAMED_BUILDER(multipleControlledRccx)})); -/// @} -/// \name QC/Operations/StandardGates/ROp.cpp -/// @{ +// QC/Operations/StandardGates/ROp.cpp INSTANTIATE_TEST_SUITE_P( QCROpTest, QCTest, testing::Values( @@ -1360,10 +1587,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledR)}, QCTestCase{"PowRScaled", MQT_NAMED_BUILDER(powRScaled), MQT_NAMED_BUILDER(powRScaledRef)})); -/// @} -/// \name QC/Operations/StandardGates/RxOp.cpp -/// @{ +// QC/Operations/StandardGates/RxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRXOpTest, QCTest, testing::Values( @@ -1385,10 +1610,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledRx)}, QCTestCase{"PowRxScaled", MQT_NAMED_BUILDER(powRxScaled), MQT_NAMED_BUILDER(rxScaled)})); -/// @} -/// \name QC/Operations/StandardGates/RxxOp.cpp -/// @{ +// QC/Operations/StandardGates/RxxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRXXOpTest, QCTest, testing::Values(QCTestCase{"RXX", MQT_NAMED_BUILDER(rxx), @@ -1410,10 +1633,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRXX", MQT_NAMED_BUILDER(inverseMultipleControlledRxx), MQT_NAMED_BUILDER(multipleControlledRxx)})); -/// @} -/// \name QC/Operations/StandardGates/RyOp.cpp -/// @{ +// QC/Operations/StandardGates/RyOp.cpp INSTANTIATE_TEST_SUITE_P( QCRYOpTest, QCTest, testing::Values( @@ -1433,10 +1654,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRY", MQT_NAMED_BUILDER(inverseMultipleControlledRy), MQT_NAMED_BUILDER(multipleControlledRy)})); -/// @} -/// \name QC/Operations/StandardGates/RyyOp.cpp -/// @{ +// QC/Operations/StandardGates/RyyOp.cpp INSTANTIATE_TEST_SUITE_P( QCRYYOpTest, QCTest, testing::Values(QCTestCase{"RYY", MQT_NAMED_BUILDER(ryy), @@ -1458,10 +1677,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRYY", MQT_NAMED_BUILDER(inverseMultipleControlledRyy), MQT_NAMED_BUILDER(multipleControlledRyy)})); -/// @} -/// \name QC/Operations/StandardGates/RzOp.cpp -/// @{ +// QC/Operations/StandardGates/RzOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZOpTest, QCTest, testing::Values( @@ -1481,10 +1698,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRZ", MQT_NAMED_BUILDER(inverseMultipleControlledRz), MQT_NAMED_BUILDER(multipleControlledRz)})); -/// @} -/// \name QC/Operations/StandardGates/RzxOp.cpp -/// @{ +// QC/Operations/StandardGates/RzxOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZXOpTest, QCTest, testing::Values(QCTestCase{"RZX", MQT_NAMED_BUILDER(rzx), @@ -1506,10 +1721,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRZX", MQT_NAMED_BUILDER(inverseMultipleControlledRzx), MQT_NAMED_BUILDER(multipleControlledRzx)})); -/// @} -/// \name QC/Operations/StandardGates/RzzOp.cpp -/// @{ +// QC/Operations/StandardGates/RzzOp.cpp INSTANTIATE_TEST_SUITE_P( QCRZZOpTest, QCTest, testing::Values(QCTestCase{"RZZ", MQT_NAMED_BUILDER(rzz), @@ -1531,10 +1744,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledRZZ", MQT_NAMED_BUILDER(inverseMultipleControlledRzz), MQT_NAMED_BUILDER(multipleControlledRzz)})); -/// @} -/// \name QC/Operations/StandardGates/SOp.cpp -/// @{ +// QC/Operations/StandardGates/SOp.cpp INSTANTIATE_TEST_SUITE_P( QCSOpTest, QCTest, testing::Values( @@ -1560,10 +1771,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(t_)}, QCTestCase{"PowThirdSToP", MQT_NAMED_BUILDER(powThirdS), MQT_NAMED_BUILDER(powThirdSRef)})); -/// @} -/// \name QC/Operations/StandardGates/SdgOp.cpp -/// @{ +// QC/Operations/StandardGates/SdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCSdgOpTest, QCTest, testing::Values(QCTestCase{"Sdg", MQT_NAMED_BUILDER(sdg), @@ -1591,10 +1800,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(tdg)}, QCTestCase{"PowThirdSdgToP", MQT_NAMED_BUILDER(powThirdSdg), MQT_NAMED_BUILDER(powThirdSdgRef)})); -/// @} -/// \name QC/Operations/StandardGates/SwapOp.cpp -/// @{ +// QC/Operations/StandardGates/SwapOp.cpp INSTANTIATE_TEST_SUITE_P( QCSWAPOpTest, QCTest, testing::Values(QCTestCase{"SWAP", MQT_NAMED_BUILDER(swap), @@ -1620,10 +1827,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(allocQubitRegister)}, QCTestCase{"PowOddSWAP", MQT_NAMED_BUILDER(powOddSwap), MQT_NAMED_BUILDER(swap)})); -/// @} -/// \name QC/Operations/StandardGates/SxOp.cpp -/// @{ +// QC/Operations/StandardGates/SxOp.cpp INSTANTIATE_TEST_SUITE_P( QCSXOpTest, QCTest, testing::Values( @@ -1647,10 +1852,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powTwoSxRef)}, QCTestCase{"PowThirdSxGeneral", MQT_NAMED_BUILDER(powThirdSx), MQT_NAMED_BUILDER(powThirdSxRef)})); -/// @} -/// \name QC/Operations/StandardGates/SxdgOp.cpp -/// @{ +// QC/Operations/StandardGates/SxdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCSXdgOpTest, QCTest, testing::Values( @@ -1676,10 +1879,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powTwoSxdgRef)}, QCTestCase{"PowThirdSxdgGeneral", MQT_NAMED_BUILDER(powThirdSxdg), MQT_NAMED_BUILDER(powThirdSxdgRef)})); -/// @} -/// \name QC/Operations/StandardGates/TOp.cpp -/// @{ +// QC/Operations/StandardGates/TOp.cpp INSTANTIATE_TEST_SUITE_P( QCTOpTest, QCTest, testing::Values( @@ -1701,10 +1902,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"PowTwoT", MQT_NAMED_BUILDER(powTwoT), MQT_NAMED_BUILDER(s)}, QCTestCase{"PowThirdTToP", MQT_NAMED_BUILDER(powThirdT), MQT_NAMED_BUILDER(powThirdTRef)})); -/// @} -/// \name QC/Operations/StandardGates/TdgOp.cpp -/// @{ +// QC/Operations/StandardGates/TdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCTdgOpTest, QCTest, testing::Values(QCTestCase{"Tdg", MQT_NAMED_BUILDER(tdg), @@ -1730,10 +1929,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sdg)}, QCTestCase{"PowThirdTdgToP", MQT_NAMED_BUILDER(powThirdTdg), MQT_NAMED_BUILDER(powThirdTdgRef)})); -/// @} -/// \name QC/Operations/StandardGates/U2Op.cpp -/// @{ +// QC/Operations/StandardGates/U2Op.cpp INSTANTIATE_TEST_SUITE_P( QCU2OpTest, QCTest, testing::Values( @@ -1753,10 +1950,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledU2", MQT_NAMED_BUILDER(inverseMultipleControlledU2), MQT_NAMED_BUILDER(multipleControlledU2)})); -/// @} -/// \name QC/Operations/StandardGates/UOp.cpp -/// @{ +// QC/Operations/StandardGates/UOp.cpp INSTANTIATE_TEST_SUITE_P( QCUOpTest, QCTest, testing::Values( @@ -1775,10 +1970,8 @@ INSTANTIATE_TEST_SUITE_P( QCTestCase{"InverseMultipleControlledU", MQT_NAMED_BUILDER(inverseMultipleControlledU), MQT_NAMED_BUILDER(multipleControlledU)})); -/// @} -/// \name QC/Operations/StandardGates/XOp.cpp -/// @{ +// QC/Operations/StandardGates/XOp.cpp INSTANTIATE_TEST_SUITE_P( QCXOpTest, QCTest, testing::Values( @@ -1803,10 +1996,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sxdg)}, QCTestCase{"PowThirdXGeneral", MQT_NAMED_BUILDER(powThirdX), MQT_NAMED_BUILDER(powThirdXRef)})); -/// @} -/// \name QC/Operations/StandardGates/XxMinusYyOp.cpp -/// @{ +// QC/Operations/StandardGates/XxMinusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCXXMinusYYOpTest, QCTest, testing::Values( @@ -1831,10 +2022,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledXxMinusYY)}, QCTestCase{"PowXxMinusYYScaled", MQT_NAMED_BUILDER(powXxMinusYYScaled), MQT_NAMED_BUILDER(powXxMinusYYScaledRef)})); -/// @} -/// \name QC/Operations/StandardGates/XxPlusYyOp.cpp -/// @{ +// QC/Operations/StandardGates/XxPlusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCXXPlusYYOpTest, QCTest, testing::Values( @@ -1859,10 +2048,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledXxPlusYY)}, QCTestCase{"PowXxPlusYYScaled", MQT_NAMED_BUILDER(powXxPlusYYScaled), MQT_NAMED_BUILDER(powXxPlusYYScaledRef)})); -/// @} -/// \name QC/Operations/StandardGates/YOp.cpp -/// @{ +// QC/Operations/StandardGates/YOp.cpp INSTANTIATE_TEST_SUITE_P( QCYOpTest, QCTest, testing::Values( @@ -1883,10 +2070,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledY)}, QCTestCase{"PowHalfY", MQT_NAMED_BUILDER(powHalfY), MQT_NAMED_BUILDER(powHalfYRef)})); -/// @} -/// \name QC/Operations/StandardGates/ZOp.cpp -/// @{ +// QC/Operations/StandardGates/ZOp.cpp INSTANTIATE_TEST_SUITE_P( QCZOpTest, QCTest, testing::Values( @@ -1911,10 +2096,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sdg)}, QCTestCase{"PowThirdZToP", MQT_NAMED_BUILDER(powThirdZ), MQT_NAMED_BUILDER(powThirdZRef)})); -/// @} -/// \name QC/QubitManagement/QubitManagement.cpp -/// @{ +// QC/QubitManagement/QubitManagement.cpp INSTANTIATE_TEST_SUITE_P( QCQubitManagementTest, QCTest, testing::Values( @@ -1944,10 +2127,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(staticQubitsCanonical)}, QCTestCase{"AllocDeallocPair", MQT_NAMED_BUILDER(allocDeallocPair), MQT_NAMED_BUILDER(emptyQC)})); -/// @} -/// \name UnrollModifiers -/// @{ +// UnrollModifiers static LogicalResult runUnrollModifiers(ModuleOp moduleOp) { PassManager pm(moduleOp.getContext()); pm.addPass(mlir::mqt::createUnrollModifiers()); @@ -2138,4 +2319,3 @@ TEST_F(QCTest, UnrollModifiersLeavesNonIntegerPowUntouched) { expectUnrollsTo(context.get(), powHalfDisjoint, powHalfDisjoint, checkPreservedPowStructure); } -/// @} diff --git a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp index c0e9a8a4f8..377bdf8b0d 100644 --- a/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp +++ b/mlir/unittests/Dialect/QCO/IR/test_qco_ir.cpp @@ -19,6 +19,7 @@ #include "mlir/Dialect/QCO/IR/QCODialect.h" #include "mlir/Dialect/QCO/IR/QCOInterfaces.h" #include "mlir/Dialect/QCO/IR/QCOOps.h" +#include "mlir/Dialect/QCO/Utils/FunctionUtils.h" #include "mlir/Dialect/QTensor/IR/QTensorDialect.h" #include "mlir/Dialect/QTensor/IR/QTensorOps.h" #include "mlir/Support/Passes.h" @@ -33,6 +34,7 @@ #include #include #include +#include #include #include #include @@ -137,17 +139,17 @@ TEST_P(QCOTest, ProgramEquivalence) { } TEST_F(QCOTest, QubitIsVectorElement) { - auto module = parseSourceString(R"mlir( + auto moduleOp = parseSourceString(R"mlir( module { func.func @f(%arg: vector<2x!qco.qubit>) { return } } )mlir", - context.get()); - ASSERT_TRUE(module); + context.get()); + ASSERT_TRUE(moduleOp); - auto function = *module->getOps().begin(); + auto function = *moduleOp->getOps().begin(); const auto vectorType = dyn_cast(function.getArgument(0).getType()); ASSERT_TRUE(vectorType); @@ -200,11 +202,11 @@ TEST_F(QCOTest, BuilderReturnsTrackedQubit) { } TEST_F(QCOTest, CleanupPreservesReturnedStaticQubit) { - auto module = QCOProgramBuilder::build( + auto moduleOp = QCOProgramBuilder::build( context.get(), [&](auto& builder) { return builder.staticQubit(0); }); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); - auto mainFunc = *module->getOps().begin(); + auto mainFunc = *moduleOp->getOps().begin(); auto returnOp = cast(mainFunc.getBody().front().back()); ASSERT_EQ(returnOp.getNumOperands(), 1U); auto returnedQubit = returnOp.getOperand(0); @@ -213,19 +215,19 @@ TEST_F(QCOTest, CleanupPreservesReturnedStaticQubit) { EXPECT_EQ(*returnedQubit.user_begin(), returnOp.getOperation()); EXPECT_TRUE(mainFunc.getBody().getOps().empty()); - ASSERT_TRUE(runQCOCleanupPipeline(*module).succeeded()); - EXPECT_TRUE(verify(*module).succeeded()); + ASSERT_TRUE(runQCOCleanupPipeline(*moduleOp).succeeded()); + EXPECT_TRUE(verify(*moduleOp).succeeded()); returnOp = cast(mainFunc.getBody().front().back()); EXPECT_TRUE(returnOp.getOperand(0).getDefiningOp()); } TEST_F(QCOTest, CleanupPreservesReturnedQubitTensor) { - auto module = QCOProgramBuilder::build( + auto moduleOp = QCOProgramBuilder::build( context.get(), [&](auto& builder) { return builder.qtensorAlloc(2); }); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); - auto mainFunc = *module->getOps().begin(); + auto mainFunc = *moduleOp->getOps().begin(); auto returnOp = cast(mainFunc.getBody().front().back()); ASSERT_EQ(returnOp.getNumOperands(), 1U); auto returnedTensor = returnOp.getOperand(0); @@ -234,8 +236,8 @@ TEST_F(QCOTest, CleanupPreservesReturnedQubitTensor) { EXPECT_EQ(*returnedTensor.user_begin(), returnOp.getOperation()); EXPECT_TRUE(mainFunc.getBody().getOps().empty()); - ASSERT_TRUE(runQCOCleanupPipeline(*module).succeeded()); - EXPECT_TRUE(verify(*module).succeeded()); + ASSERT_TRUE(runQCOCleanupPipeline(*moduleOp).succeeded()); + EXPECT_TRUE(verify(*moduleOp).succeeded()); returnOp = cast(mainFunc.getBody().front().back()); EXPECT_TRUE(returnOp.getOperand(0).getDefiningOp()); @@ -323,6 +325,217 @@ TEST_F(QCOTest, BuilderSupportsIndependentClassicalRegisterInitialization) { "undefined"); } +TEST_F(QCOTest, BuilderCreatesGenericAndUnitaryFunctions) { + QCOProgramBuilder builder(context.get()); + builder.initialize(); + auto qubitType = QubitType::get(context.get()); + + auto reset = builder.createFunction( + "reset", TypeRange{qubitType}, [&](ValueRange arguments) { + return SmallVector{builder.reset(arguments[0])}; + }); + auto flip = builder.createUnitaryFunction( + "flip", TypeRange{qubitType}, [&](ValueRange arguments) { + return SmallVector{builder.x(arguments[0])}; + }); + + Value qubit = builder.allocQubit(); + qubit = builder.call(reset, qubit).front(); + qubit = builder.call(flip, qubit).front(); + qubit = builder.inv(qubit, [&](Value argument) { + return builder.call(flip, argument).front(); + }); + builder.sink(qubit); + auto moduleOp = builder.finalize(); + + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + EXPECT_EQ(reset.getResultTypes(), reset.getArgumentTypes()); + EXPECT_EQ(flip.getResultTypes(), flip.getArgumentTypes()); + EXPECT_FALSE(mlir::mqt::isUnitaryFunction(reset)); + EXPECT_TRUE(mlir::mqt::isUnitaryFunction(flip)); + + auto mainFunc = mlir::mqt::getEntryPoint(*moduleOp); + ASSERT_TRUE(mainFunc); + EXPECT_EQ(llvm::range_size(mainFunc.getBody().getOps()), 1U); + EXPECT_EQ(llvm::range_size(mainFunc.getBody().getOps()), 1U); + auto call = *mainFunc.getBody().getOps().begin(); + EXPECT_FALSE(call.getInputForOutput(qubit)); + EXPECT_FALSE(call.getOutputForInput(qubit)); + auto inverse = *mainFunc.getBody().getOps().begin(); + EXPECT_TRUE(isa(&inverse.getRegion().front().front())); +} + +TEST_F(QCOTest, BuilderFinalizesRenamedEntryPoint) { + QCOProgramBuilder builder(context.get()); + builder.initialize(); + auto entry = cast(builder.getInsertionBlock()->getParentOp()); + entry.setName("entry"); + + auto moduleOp = builder.finalize(); + + ASSERT_TRUE(moduleOp); + EXPECT_EQ(mlir::mqt::getEntryPoint(*moduleOp).getName(), "entry"); +} + +TEST_F(QCOTest, UnitaryVerifierDiagnosesMalformedCalls) { + ParserConfig config(context.get(), false); + auto moduleOp = parseSourceString(R"mlir( + module { + func.func private @malformed(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %left, %right = qco.call @malformed(%q) + : (!qco.qubit) -> (!qco.qubit, !qco.qubit) + return %left : !qco.qubit + } + } + )mlir", + config); + ASSERT_TRUE(moduleOp); + + bool sawExpectedDiagnostic = false; + ScopedDiagnosticHandler handler(context.get(), [&](Diagnostic& diagnostic) { + sawExpectedDiagnostic |= + StringRef(diagnostic.str()) + .contains("requires one trailing qubit operand for every qubit " + "result"); + return success(); + }); + EXPECT_TRUE(failed(verify(*moduleOp))); + EXPECT_TRUE(sawExpectedDiagnostic); +} + +TEST_F(QCOTest, UnitaryVerifierRejectsInvalidFunctionAndCallContracts) { + DialectRegistry registry; + registry.insert(); + context->appendDialectRegistry(registry); + context->getOrLoadDialect(); + + constexpr std::array invalidPrograms{ + R"mlir(module { + func.func private @bad() attributes {mqt.unitary} { return } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qco.qubit, %theta: f64) + -> !qco.qubit attributes {mqt.unitary} { + return %q : !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qco.qubit) + attributes {mqt.unitary} { return } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.reset %q : !qco.qubit -> !qco.qubit + return %out : !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @bad(%left: !qco.qubit, %right: !qco.qubit) + -> (!qco.qubit, !qco.qubit) attributes {mqt.unitary} { + return %right, %left : !qco.qubit, !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.call @bad(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @bad(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.call @missing(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @plain(%q: !qco.qubit) -> !qco.qubit { + return %q : !qco.qubit + } + func.func @main(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.entry_point} { + %out = qco.call @plain(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + })mlir", + R"mlir(module { + func.func private @flip(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.unitary} { + %out = qco.x %q : !qco.qubit -> !qco.qubit + return %out : !qco.qubit + } + func.func @main(%left: !qco.qubit, %right: !qco.qubit) + -> (!qco.qubit, !qco.qubit) attributes {mqt.entry_point} { + %a, %b = qco.call @flip(%left, %right) + : (!qco.qubit, !qco.qubit) -> (!qco.qubit, !qco.qubit) + return %a, %b : !qco.qubit, !qco.qubit + } + })mlir", + }; + + ParserConfig config(context.get(), false); + for (const auto source : invalidPrograms) { + auto moduleOp = parseSourceString(source, config); + ASSERT_TRUE(moduleOp); + EXPECT_TRUE(failed(verify(*moduleOp))); + } + + auto resultModule = parseSourceString(R"mlir(module { + func.func private @bad(%q: !qco.qubit) + -> (!qco.qubit, !qco.qubit) attributes {mqt.unitary} + func.func @main(%q: !qco.qubit) -> !qco.qubit + attributes {mqt.entry_point} { + %out = qco.call @bad(%q) : (!qco.qubit) -> !qco.qubit + return %out : !qco.qubit + } + })mlir", + config); + ASSERT_TRUE(resultModule); + auto call = *mlir::mqt::getEntryPoint(*resultModule) + .getBody() + .getOps() + .begin(); + SymbolTableCollection symbols; + EXPECT_TRUE(failed(call.verifySymbolUses(symbols))); +} + +TEST_F(QCOTest, TraceQubitArgumentRejectsUnsupportedSources) { + ParserConfig config(context.get(), false); + auto moduleOp = parseSourceString(R"mlir(module { + func.func private @declaration(!qco.qubit) -> !qco.qubit + func.func private @callee(%q: !qco.qubit) -> (i1, !qco.qubit) { + %flag = arith.constant true + return %flag, %q : i1, !qco.qubit + } + func.func @main(%q: !qco.qubit) -> (i1, !qco.qubit) { + %flag, %out = func.call @callee(%q) + : (!qco.qubit) -> (i1, !qco.qubit) + %missing = func.call @missing(%q) : (!qco.qubit) -> !qco.qubit + %constant = arith.constant true + return %flag, %out : i1, !qco.qubit + } + })mlir", + config); + ASSERT_TRUE(moduleOp); + auto declaration = moduleOp->lookupSymbol("declaration"); + auto callee = moduleOp->lookupSymbol("callee"); + auto main = moduleOp->lookupSymbol("main"); + ASSERT_TRUE(declaration && callee && main); + auto calls = llvm::to_vector(main.getOps()); + ASSERT_EQ(calls.size(), 2U); + auto constant = *main.getOps().begin(); + + EXPECT_TRUE(failed(traceQubitArgument(declaration, {}))); + EXPECT_TRUE(failed(traceQubitArgument(main, callee.getArgument(0)))); + EXPECT_TRUE(failed(traceQubitArgument(main, calls[0].getResult(0)))); + EXPECT_TRUE(failed(traceQubitArgument(main, calls[1].getResult(0)))); + EXPECT_TRUE(failed(traceQubitArgument(main, constant.getResult()))); +} + TEST_F(QCOTest, DirectSingleQubitPowBuilder) { QCOProgramBuilder builder(context.get()); builder.initialize(); @@ -790,12 +1003,12 @@ TEST_F(QCOTest, IfOpWithClassicalResultRoundTripsAndPreservesTies) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); IfOp ifOp; - module->walk([&](IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); ASSERT_TRUE(ifOp); ASSERT_EQ(ifOp.getClassicalResults().size(), 1); ASSERT_EQ(ifOp.getLinearResults().size(), 1); @@ -831,13 +1044,13 @@ TEST_F(QCOTest, IfOpWithClassicalResultRoundTripsAndPreservesTies) { std::string printed; llvm::raw_string_ostream stream(printed); - module->print(stream); + moduleOp->print(stream); stream.flush(); auto reparsedModule = parseSourceString(printed, context.get()); ASSERT_TRUE(reparsedModule); EXPECT_TRUE(succeeded(verify(*reparsedModule))); - EXPECT_TRUE( - areModulesEquivalentWithPermutations(module.get(), reparsedModule.get())); + EXPECT_TRUE(areModulesEquivalentWithPermutations(moduleOp.get(), + reparsedModule.get())); } TEST_F(QCOTest, IfOpRejectsMismatchedClassicalYield) { @@ -917,16 +1130,16 @@ TEST_F(QCOTest, CanonicalizesConstantIfWithClassicalResult) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(module.get()))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); bool containsIf = false; - module->walk([&](IfOp) { containsIf = true; }); + moduleOp->walk([&](IfOp) { containsIf = true; }); EXPECT_FALSE(containsIf); - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); APInt result; @@ -960,14 +1173,14 @@ TEST_F(QCOTest, CanonicalizesRedundantClassicalIfResults) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(module.get()))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); IfOp ifOp; - module->walk([&](IfOp candidate) { ifOp = candidate; }); + moduleOp->walk([&](IfOp candidate) { ifOp = candidate; }); ASSERT_TRUE(ifOp); ASSERT_EQ(ifOp.getClassicalResults().size(), 1); ASSERT_EQ(ifOp.getLinearResults().size(), 1); @@ -981,7 +1194,7 @@ TEST_F(QCOTest, CanonicalizesRedundantClassicalIfResults) { ifOp.getLinearResults().front().getType()); } - auto main = module->lookupSymbol("main"); + auto main = moduleOp->lookupSymbol("main"); ASSERT_TRUE(main); auto returnOp = cast(main.getBody().front().getTerminator()); ASSERT_EQ(returnOp.getNumOperands(), 3); @@ -1444,12 +1657,12 @@ TEST_F(QCOTest, IndexSwitchWithClassicalResultRoundTripsAndPreservesTies) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); IndexSwitchOp switchOp; - module->walk([&](IndexSwitchOp candidate) { switchOp = candidate; }); + moduleOp->walk([&](IndexSwitchOp candidate) { switchOp = candidate; }); ASSERT_TRUE(switchOp); ASSERT_EQ(switchOp.getClassicalResults().size(), 1); ASSERT_EQ(switchOp.getLinearResults().size(), 1); @@ -1499,13 +1712,13 @@ TEST_F(QCOTest, IndexSwitchWithClassicalResultRoundTripsAndPreservesTies) { std::string printed; llvm::raw_string_ostream stream(printed); - module->print(stream); + moduleOp->print(stream); stream.flush(); auto reparsedModule = parseSourceString(printed, context.get()); ASSERT_TRUE(reparsedModule); EXPECT_TRUE(succeeded(verify(*reparsedModule))); - EXPECT_TRUE( - areModulesEquivalentWithPermutations(module.get(), reparsedModule.get())); + EXPECT_TRUE(areModulesEquivalentWithPermutations(moduleOp.get(), + reparsedModule.get())); } TEST_F(QCOTest, ClassicalYieldOrderAffectsConditionalEquivalence) { @@ -1556,9 +1769,9 @@ TEST_F(QCOTest, ClassicalYieldOrderAffectsConditionalEquivalence) { ASSERT_TRUE(lhs); ASSERT_TRUE(rhs); - const auto findFirstYield = [](ModuleOp module) { + const auto findFirstYield = [](ModuleOp moduleOp) { YieldOp result; - module.walk([&](YieldOp candidate) { + moduleOp.walk([&](YieldOp candidate) { if (!result) { result = candidate; } @@ -1580,14 +1793,14 @@ TEST_F(QCOTest, ClassicalYieldOrderAffectsConditionalEquivalence) { auto duplicateRhs = parseSourceString(source, context.get()); ASSERT_TRUE(duplicateLhs); ASSERT_TRUE(duplicateRhs); - for (ModuleOp module : {*duplicateLhs, *duplicateRhs}) { - auto yield = findFirstYield(module); + for (ModuleOp moduleOp : {*duplicateLhs, *duplicateRhs}) { + auto yield = findFirstYield(moduleOp); ASSERT_TRUE(yield); SmallVector duplicateOperands(yield.getTargets()); ASSERT_GE(duplicateOperands.size(), 2); duplicateOperands[1] = duplicateOperands[0]; yield->setOperands(duplicateOperands); - ASSERT_TRUE(succeeded(verify(module))); + ASSERT_TRUE(succeeded(verify(moduleOp))); } EXPECT_TRUE(areModulesEquivalentWithPermutations(duplicateLhs.get(), duplicateRhs.get())); @@ -1614,10 +1827,10 @@ TEST_F(QCOTest, ExtendsMixedResultIndexSwitchTargets) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); IndexSwitchOp switchOp; - module->walk([&](IndexSwitchOp candidate) { switchOp = candidate; }); + moduleOp->walk([&](IndexSwitchOp candidate) { switchOp = candidate; }); ASSERT_TRUE(switchOp); IRRewriter rewriter(context.get()); @@ -1629,7 +1842,7 @@ TEST_F(QCOTest, ExtendsMixedResultIndexSwitchTargets) { SinkOp::create(rewriter, extended.getLoc(), extended.getLinearResults().back()); - ASSERT_TRUE(succeeded(verify(*module))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); ASSERT_EQ(extended.getClassicalResults().size(), 1); ASSERT_EQ(extended.getLinearResults().size(), 2); for (Region* region : extended.getRegions()) { @@ -1729,7 +1942,7 @@ TEST_F(QCOTest, IndexSwitchConstantSuccessor) { auto result = builder.qcoIndexSwitch(1, q0, SmallVector{0, 1}, caseBodies, identity); builder.sink(result); - [[maybe_unused]] auto module = builder.finalize(); + [[maybe_unused]] auto moduleOp = builder.finalize(); auto switchOp = result.getDefiningOp(); ASSERT_TRUE(switchOp); @@ -1822,20 +2035,20 @@ TEST_F(QCOTest, CanonicalizesConstantIndexSwitchToSelectedCaseOrDefault) { } )mlir"; - auto module = parseSourceString(mlirCode, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); - ASSERT_TRUE(succeeded(runQCOCleanupPipeline(module.get()))); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(mlirCode, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + ASSERT_TRUE(succeeded(runQCOCleanupPipeline(moduleOp.get()))); + ASSERT_TRUE(succeeded(verify(*moduleOp))); bool containsSwitch = false; - module->walk([&](IndexSwitchOp) { containsSwitch = true; }); + moduleOp->walk([&](IndexSwitchOp) { containsSwitch = true; }); EXPECT_FALSE(containsSwitch); const auto checkSelectedRegion = [&](const StringRef functionName, const int64_t expectedNumber, const StringRef expectedGate) { - auto func = module->lookupSymbol(functionName); + auto func = moduleOp->lookupSymbol(functionName); ASSERT_TRUE(func); HOp consumer; @@ -1857,8 +2070,7 @@ TEST_F(QCOTest, CanonicalizesConstantIndexSwitchToSelectedCaseOrDefault) { checkSelectedRegion("selected_default", 22, "qco.z"); } -/// \name QCO/SCF/IfOp.cpp -/// @{ +// QCO/SCF/IfOp.cpp INSTANTIATE_TEST_SUITE_P( QCOIfOpTest, QCOTest, testing::Values( @@ -1878,10 +2090,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(simpleIf)}, QCOTestCase{"NestedFalseIf", MQT_NAMED_BUILDER(nestedFalseIf), MQT_NAMED_BUILDER(ifElse)})); -/// @} -/// \name QCO/Modifiers/CtrlOp.cpp -/// @{ +// QCO/Modifiers/CtrlOp.cpp INSTANTIATE_TEST_SUITE_P( QCOCtrlOpTest, QCOTest, testing::Values( @@ -1903,10 +2113,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"ModifierBodyReuseReordered", MQT_NAMED_BUILDER(modifierBodyReuseReordered), MQT_NAMED_BUILDER(modifierBodyReuseReorderedRef)})); -/// @} -/// \name QCO/Modifiers/InvOp.cpp -/// @{ +// QCO/Modifiers/InvOp.cpp INSTANTIATE_TEST_SUITE_P( QCOInvOpTest, QCOTest, testing::Values(QCOTestCase{"EmptyInv", MQT_NAMED_BUILDER(emptyInv), @@ -1923,7 +2131,6 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(ctrlInvTwo)}, QCOTestCase{"InverseT", MQT_NAMED_BUILDER(inverseT), MQT_NAMED_BUILDER(tdg)})); -/// @} /// A power modifier with a qubit that its body does not use. static Value powWithUnusedQubit(QCOProgramBuilder& b) { @@ -1935,8 +2142,7 @@ static Value powWithUnusedQubit(QCOProgramBuilder& b) { return measureRegister(b, powOut); } -/// \name QCO/Modifiers/PowOp.cpp -/// @{ +// QCO/Modifiers/PowOp.cpp INSTANTIATE_TEST_SUITE_P( QCOPowOpTest, QCOTest, testing::Values( @@ -1973,7 +2179,6 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc1QubitRegister)}, QCOTestCase{"PowWithUnusedQubit", MQT_NAMED_BUILDER(powWithUnusedQubit), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} TEST_F(QCOTest, PowExponentIsUnitaryParameter) { auto program = @@ -2068,8 +2273,8 @@ TEST_F(QCOTest, NestedPowAcrossBranchCutDoesNotMerge) { EXPECT_TRUE(matrix->isApprox(DynamicMatrix::identity(2), 1e-10)); } -/// pow(rxx) folds the exponent into the rotation angle: pow(2){rxx(θ)} => -/// rxx(2θ). Verify cleanup and the hoisted parameter's SSA dominance. +// pow(rxx) folds the exponent into the rotation angle: pow(2){rxx(θ)} => +// rxx(2θ). Verify cleanup and the hoisted parameter's SSA dominance. TEST_F(QCOTest, PowRxxFold) { auto program = ::mqt::test::buildMLIRProgram(context.get(), MQT_NAMED_BUILDER(powRxx)); @@ -2196,8 +2401,8 @@ TEST_F(QCOTest, EvenPowFoldPreservesReorderedBodyResults) { EXPECT_EQ(measurements[1].getQubitIn(), allocations[0].getResult()); } -/// pow(-0.5) { h } cannot fold a negative fractional exponent -/// into H (no angle to scale). Verify that PowOp survives. +// pow(-0.5) { h } cannot fold a negative fractional exponent +// into H (no angle to scale). Verify that PowOp survives. TEST_F(QCOTest, NegPowHNoFold) { auto program = ::mqt::test::buildMLIRProgram(context.get(), MQT_NAMED_BUILDER(negPowH)); @@ -2211,9 +2416,9 @@ TEST_F(QCOTest, NegPowHNoFold) { EXPECT_EQ(powCount, 1) << "PowOp around h must survive the pipeline"; } -/// pow(sx) inside a ctrl modifier expands into GPhase + RX. Global-phase -/// normalization then turns the controlled GPhase into P on the control. -/// Verify the CtrlOp survives and the relative phase remains observable. +// pow(sx) inside a ctrl modifier expands into GPhase + RX. Global-phase +// normalization then turns the controlled GPhase into P on the control. +// Verify the CtrlOp survives and the relative phase remains observable. TEST_F(QCOTest, CtrlPowSxExpands) { auto program = ::mqt::test::buildMLIRProgram(context.get(), MQT_NAMED_BUILDER(ctrlPowSx)); @@ -2263,8 +2468,7 @@ TEST_F(QCOTest, CtrlGPhasePassesTargetsThrough) { EXPECT_TRUE(mainFunc.getBody().getOps().empty()); } -/// \name QCO/Operations/StandardGates/BarrierOp.cpp -/// @{ +// QCO/Operations/StandardGates/BarrierOp.cpp INSTANTIATE_TEST_SUITE_P( QCOBarrierOpTest, QCOTest, testing::Values(QCOTestCase{"Barrier", MQT_NAMED_BUILDER(barrier), @@ -2285,10 +2489,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(barrierTwoQubits)}, QCOTestCase{"PowBarrier", MQT_NAMED_BUILDER(powBarrier), MQT_NAMED_BUILDER(barrier)})); -/// @} -/// \name QCO/Operations/StandardGates/DcxOp.cpp -/// @{ +// QCO/Operations/StandardGates/DcxOp.cpp INSTANTIATE_TEST_SUITE_P( QCODCXOpTest, QCOTest, testing::Values( @@ -2315,10 +2517,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoDCXSwappedTargets", MQT_NAMED_BUILDER(twoDcxSwappedTargets), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/EcrOp.cpp -/// @{ +// QCO/Operations/StandardGates/EcrOp.cpp INSTANTIATE_TEST_SUITE_P( QCOECROpTest, QCOTest, testing::Values(QCOTestCase{"ECR", MQT_NAMED_BUILDER(ecr), @@ -2346,10 +2546,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc2QubitRegister)}, QCOTestCase{"PowOddECR", MQT_NAMED_BUILDER(powOddEcr), MQT_NAMED_BUILDER(ecr)})); -/// @} -/// \name QCO/Operations/StandardGates/GphaseOp.cpp -/// @{ +// QCO/Operations/StandardGates/GphaseOp.cpp INSTANTIATE_TEST_SUITE_P( QCOGPhaseOpTest, QCOTest, testing::Values( @@ -2370,10 +2568,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powGphaseScaledRef)}, QCOTestCase{"NegPowGphase", MQT_NAMED_BUILDER(negPowGphase), MQT_NAMED_BUILDER(negPowGphaseRef)})); -/// @} -/// \name QCO/Operations/StandardGates/HOp.cpp -/// @{ +// QCO/Operations/StandardGates/HOp.cpp INSTANTIATE_TEST_SUITE_P( QCOHOpTest, QCOTest, testing::Values( @@ -2398,10 +2594,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc1QubitRegister)}, QCOTestCase{"PowOddH", MQT_NAMED_BUILDER(powOddH), MQT_NAMED_BUILDER(h)})); -/// @} -/// \name QCO/Operations/StandardGates/IdOp.cpp -/// @{ +// QCO/Operations/StandardGates/IdOp.cpp INSTANTIATE_TEST_SUITE_P( QCOIDOpTest, QCOTest, testing::Values( @@ -2426,10 +2620,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc3QubitRegister)}, QCOTestCase{"PowId", MQT_NAMED_BUILDER(powId), MQT_NAMED_BUILDER(alloc1QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/IswapOp.cpp -/// @{ +// QCO/Operations/StandardGates/IswapOp.cpp INSTANTIATE_TEST_SUITE_P( QCOiSWAPOpTest, QCOTest, testing::Values(QCOTestCase{"iSWAP", MQT_NAMED_BUILDER(iswap), @@ -2454,10 +2646,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(inverseMultipleControlledIswap)}, QCOTestCase{"PowHalfiSWAP", MQT_NAMED_BUILDER(powHalfIswap), MQT_NAMED_BUILDER(powHalfIswapRef)})); -/// @} -/// \name QCO/Operations/StandardGates/POp.cpp -/// @{ +// QCO/Operations/StandardGates/POp.cpp INSTANTIATE_TEST_SUITE_P( QCOPOpTest, QCOTest, testing::Values( @@ -2478,10 +2668,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledP)}, QCOTestCase{"TwoPOppositePhase", MQT_NAMED_BUILDER(twoPOppositePhase), MQT_NAMED_BUILDER(allocQubit)})); -/// @} -/// \name QCO/Operations/StandardGates/RCCXOp.cpp -/// @{ +// QCO/Operations/StandardGates/RCCXOp.cpp INSTANTIATE_TEST_SUITE_P( QCORCCXOpTest, QCOTest, testing::Values( @@ -2509,10 +2697,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledRccx)}, QCOTestCase{"TwoRCCX", MQT_NAMED_BUILDER(twoRccx), MQT_NAMED_BUILDER(alloc3QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/ROp.cpp -/// @{ +// QCO/Operations/StandardGates/ROp.cpp INSTANTIATE_TEST_SUITE_P( QCOROpTest, QCOTest, testing::Values( @@ -2538,10 +2724,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoR", MQT_NAMED_BUILDER(twoR), MQT_NAMED_BUILDER(r)}, QCOTestCase{"PowRScaled", MQT_NAMED_BUILDER(powRScaled), MQT_NAMED_BUILDER(powRScaledRef)})); -/// @} -/// \name QCO/Operations/StandardGates/RxOp.cpp -/// @{ +// QCO/Operations/StandardGates/RxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORXOpTest, QCOTest, testing::Values( @@ -2565,10 +2749,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc1QubitRegister)}, QCOTestCase{"PowRxScaled", MQT_NAMED_BUILDER(powRxScaled), MQT_NAMED_BUILDER(rxScaled)})); -/// @} -/// \name QCO/Operations/StandardGates/RxxOp.cpp -/// @{ +// QCO/Operations/StandardGates/RxxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORXXOpTest, QCOTest, testing::Values( @@ -2601,10 +2783,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoRXXOppositePhaseSwappedTargets", MQT_NAMED_BUILDER(twoRxxOppositePhaseSwappedTargets), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/RyOp.cpp -/// @{ +// QCO/Operations/StandardGates/RyOp.cpp INSTANTIATE_TEST_SUITE_P( QCORYOpTest, QCOTest, testing::Values( @@ -2626,10 +2806,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledRy)}, QCOTestCase{"TwoRYOppositePhase", MQT_NAMED_BUILDER(twoRyOppositePhase), MQT_NAMED_BUILDER(alloc1QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/RyyOp.cpp -/// @{ +// QCO/Operations/StandardGates/RyyOp.cpp INSTANTIATE_TEST_SUITE_P( QCORYYOpTest, QCOTest, testing::Values( @@ -2662,10 +2840,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoRYYOppositePhase", MQT_NAMED_BUILDER(twoRyyOppositePhase), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/RzOp.cpp -/// @{ +// QCO/Operations/StandardGates/RzOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZOpTest, QCOTest, testing::Values( @@ -2687,10 +2863,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(multipleControlledRz)}, QCOTestCase{"TwoRZOppositePhase", MQT_NAMED_BUILDER(twoRzOppositePhase), MQT_NAMED_BUILDER(alloc1QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/RzxOp.cpp -/// @{ +// QCO/Operations/StandardGates/RzxOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZXOpTest, QCOTest, testing::Values(QCOTestCase{"RZX", MQT_NAMED_BUILDER(rzx), @@ -2715,10 +2889,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoRZXOppositePhase", MQT_NAMED_BUILDER(twoRzxOppositePhase), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/RzzOp.cpp -/// @{ +// QCO/Operations/StandardGates/RzzOp.cpp INSTANTIATE_TEST_SUITE_P( QCORZZOpTest, QCOTest, testing::Values( @@ -2751,10 +2923,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"TwoRZZOppositePhase", MQT_NAMED_BUILDER(twoRzzOppositePhase), MQT_NAMED_BUILDER(alloc2QubitRegister)})); -/// @} -/// \name QCO/Operations/StandardGates/SOp.cpp -/// @{ +// QCO/Operations/StandardGates/SOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSOpTest, QCOTest, testing::Values( @@ -2784,10 +2954,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(t_)}, QCOTestCase{"PowThirdSToP", MQT_NAMED_BUILDER(powThirdS), MQT_NAMED_BUILDER(powThirdSRef)})); -/// @} -/// \name QCO/Operations/StandardGates/SdgOp.cpp -/// @{ +// QCO/Operations/StandardGates/SdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSdgOpTest, QCOTest, testing::Values( @@ -2818,10 +2986,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(tdg)}, QCOTestCase{"PowThirdSdgToP", MQT_NAMED_BUILDER(powThirdSdg), MQT_NAMED_BUILDER(powThirdSdgRef)})); -/// @} -/// \name QCO/Operations/StandardGates/SwapOp.cpp -/// @{ +// QCO/Operations/StandardGates/SwapOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSWAPOpTest, QCOTest, testing::Values( @@ -2852,10 +3018,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc2QubitRegister)}, QCOTestCase{"PowOddSWAP", MQT_NAMED_BUILDER(powOddSwap), MQT_NAMED_BUILDER(swap)})); -/// @} -/// \name QCO/Operations/StandardGates/SxOp.cpp -/// @{ +// QCO/Operations/StandardGates/SxOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSXOpTest, QCOTest, testing::Values( @@ -2882,10 +3046,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powTwoSxRef)}, QCOTestCase{"PowThirdSxGeneral", MQT_NAMED_BUILDER(powThirdSx), MQT_NAMED_BUILDER(powThirdSxRef)})); -/// @} -/// \name QCO/Operations/StandardGates/SxdgOp.cpp -/// @{ +// QCO/Operations/StandardGates/SxdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOSXdgOpTest, QCOTest, testing::Values( @@ -2915,10 +3077,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(powTwoSxdgRef)}, QCOTestCase{"PowThirdSxdgGeneral", MQT_NAMED_BUILDER(powThirdSxdg), MQT_NAMED_BUILDER(powThirdSxdgRef)})); -/// @} -/// \name QCO/Operations/StandardGates/TOp.cpp -/// @{ +// QCO/Operations/StandardGates/TOp.cpp INSTANTIATE_TEST_SUITE_P( QCOTOpTest, QCOTest, testing::Values( @@ -2944,10 +3104,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(s)}, QCOTestCase{"PowThirdTToP", MQT_NAMED_BUILDER(powThirdT), MQT_NAMED_BUILDER(powThirdTRef)})); -/// @} -/// \name QCO/Operations/StandardGates/TdgOp.cpp -/// @{ +// QCO/Operations/StandardGates/TdgOp.cpp INSTANTIATE_TEST_SUITE_P( QCOTdgOpTest, QCOTest, testing::Values( @@ -2977,10 +3135,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sdg)}, QCOTestCase{"PowThirdTdgToP", MQT_NAMED_BUILDER(powThirdTdg), MQT_NAMED_BUILDER(powThirdTdgRef)})); -/// @} -/// \name QCO/Operations/StandardGates/U2Op.cpp -/// @{ +// QCO/Operations/StandardGates/U2Op.cpp INSTANTIATE_TEST_SUITE_P( QCOU2OpTest, QCOTest, testing::Values( @@ -3006,10 +3162,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(rxPiOver2)}, QCOTestCase{"CanonicalizeU2ToRy", MQT_NAMED_BUILDER(canonicalizeU2ToRy), MQT_NAMED_BUILDER(ryPiOver2)})); -/// @} -/// \name QCO/Operations/StandardGates/UOp.cpp -/// @{ +// QCO/Operations/StandardGates/UOp.cpp INSTANTIATE_TEST_SUITE_P( QCOUOpTest, QCOTest, testing::Values( @@ -3036,10 +3190,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(ry)}, QCOTestCase{"CanonicalizeUToU2", MQT_NAMED_BUILDER(canonicalizeUToU2), MQT_NAMED_BUILDER(u2)})); -/// @} -/// \name QCO/Operations/StandardGates/XOp.cpp -/// @{ +// QCO/Operations/StandardGates/XOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXOpTest, QCOTest, testing::Values( @@ -3070,10 +3222,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sxdg)}, QCOTestCase{"PowThirdXGeneral", MQT_NAMED_BUILDER(powThirdX), MQT_NAMED_BUILDER(powThirdXRef)})); -/// @} -/// \name QCO/Operations/StandardGates/XxMinusYyOp.cpp -/// @{ +// QCO/Operations/StandardGates/XxMinusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXXMinusYYOpTest, QCOTest, testing::Values( @@ -3104,10 +3254,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(xxMinusYY)}, QCOTestCase{"PowXxMinusYYScaled", MQT_NAMED_BUILDER(powXxMinusYYScaled), MQT_NAMED_BUILDER(powXxMinusYYScaledRef)})); -/// @} -/// \name QCO/Operations/StandardGates/XxPlusYyOp.cpp -/// @{ +// QCO/Operations/StandardGates/XxPlusYyOp.cpp INSTANTIATE_TEST_SUITE_P( QCOXXPlusYYOpTest, QCOTest, testing::Values( @@ -3138,10 +3286,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(xxPlusYY)}, QCOTestCase{"PowXxPlusYYScaled", MQT_NAMED_BUILDER(powXxPlusYYScaled), MQT_NAMED_BUILDER(powXxPlusYYScaledRef)})); -/// @} -/// \name QCO/Operations/StandardGates/YOp.cpp -/// @{ +// QCO/Operations/StandardGates/YOp.cpp INSTANTIATE_TEST_SUITE_P( QCOYOpTest, QCOTest, testing::Values( @@ -3164,10 +3310,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(alloc1QubitRegister)}, QCOTestCase{"PowHalfY", MQT_NAMED_BUILDER(powHalfY), MQT_NAMED_BUILDER(powHalfYRef)})); -/// @} -/// \name QCO/Operations/StandardGates/ZOp.cpp -/// @{ +// QCO/Operations/StandardGates/ZOp.cpp INSTANTIATE_TEST_SUITE_P( QCOZOpTest, QCOTest, testing::Values( @@ -3194,10 +3338,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(sdg)}, QCOTestCase{"PowThirdZToP", MQT_NAMED_BUILDER(powThirdZ), MQT_NAMED_BUILDER(powThirdZRef)})); -/// @} -/// \name QCO/Operations/MeasureOp.cpp -/// @{ +// QCO/Operations/MeasureOp.cpp INSTANTIATE_TEST_SUITE_P( QCOMeasureOpTest, QCOTest, testing::Values( @@ -3214,10 +3356,8 @@ INSTANTIATE_TEST_SUITE_P( "MultipleClassicalRegistersAndMeasurements", MQT_NAMED_BUILDER(multipleClassicalRegistersAndMeasurements), MQT_NAMED_BUILDER(multipleClassicalRegistersAndMeasurements)})); -/// @} -/// \name QCO/Operations/ResetOp.cpp -/// @{ +// QCO/Operations/ResetOp.cpp INSTANTIATE_TEST_SUITE_P( QCOResetOpTest, QCOTest, testing::Values(QCOTestCase{"ResetQubitWithoutOp", @@ -3239,10 +3379,8 @@ INSTANTIATE_TEST_SUITE_P( QCOTestCase{"RepeatedResetAfterSingleOp", MQT_NAMED_BUILDER(repeatedResetAfterSingleOp), MQT_NAMED_BUILDER(resetQubitAfterSingleOp)})); -/// @} -/// \name QCO/QubitManagement/QubitManagement.cpp -/// @{ +// QCO/QubitManagement/QubitManagement.cpp INSTANTIATE_TEST_SUITE_P( QCOQubitManagementTest, QCOTest, testing::Values( @@ -3268,10 +3406,8 @@ INSTANTIATE_TEST_SUITE_P( MQT_NAMED_BUILDER(staticQubitsWithInv)}, QCOTestCase{"AllocSinkPair", MQT_NAMED_BUILDER(allocSinkPair), MQT_NAMED_BUILDER(allocQubitNoMeasure)})); -/// @} -/// \name UnrollModifiers -/// @{ +// UnrollModifiers static LogicalResult runUnrollModifiers(ModuleOp moduleOp) { PassManager pm(moduleOp.getContext()); pm.addPass(mlir::mqt::createUnrollModifiers()); @@ -3505,4 +3641,3 @@ TEST_F(QCOTest, UnrollModifiersLeavesNonIntegerPowUntouched) { expectUnrollsTo(context.get(), powHalfDisjoint, powHalfDisjoint, checkPreservedPowStructure); } -/// @} diff --git a/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp b/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp index 92d7fb4f63..381b773825 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp @@ -93,7 +93,7 @@ TEST_F(WireIteratorFixture, TraversalRespectsStraightLineSemantics) { const auto [q12, c1] = builder.measure(q11); builder.sink(q03); builder.sink(q12); - [[maybe_unused]] auto module = builder.finalize(c0); + [[maybe_unused]] auto moduleOp = builder.finalize(c0); const auto fwChain0 = getChain(q00); const auto fwChain1 = getChain(q10); @@ -127,7 +127,7 @@ TEST_F(WireIteratorFixture, TraversalVisitsSourcesAndSinks) { const auto q0 = builder.staticQubit(0); builder.sink(q0); - [[maybe_unused]] auto module = builder.finalize(); + [[maybe_unused]] auto moduleOp = builder.finalize(); WireIterator it(q0); ASSERT_EQ(it.qubit(), q0); @@ -164,7 +164,7 @@ TEST_F(WireIteratorFixture, TraversalRespectsNestedBoundaries) { })[0]; builder.sink(q1); - [[maybe_unused]] auto module = builder.finalize(); + [[maybe_unused]] auto moduleOp = builder.finalize(); WireIterator it(inLoop); ASSERT_EQ(it.qubit(), inLoop); @@ -195,7 +195,7 @@ TEST_F(WireIteratorFixture, FailOnSentinelAccess) { const auto q0 = builder.staticQubit(0); builder.sink(q0); - [[maybe_unused]] auto module = builder.finalize(); + [[maybe_unused]] auto moduleOp = builder.finalize(); WireIterator it(q0); --it; @@ -248,7 +248,7 @@ TEST_F(WireIteratorFixture, TraversalRespectsStructuredSemantics) { const auto tensor2 = builder.qtensorInsert(q15, tensor1, 1); builder.qtensorDealloc(tensor2); - [[maybe_unused]] auto module = builder.finalize(); + [[maybe_unused]] auto moduleOp = builder.finalize(); const auto fwChain0 = getChain(q00); const auto fwChain1 = getChain(q10); @@ -288,13 +288,13 @@ TEST_F(WireIteratorFixture, TraversalRespectsStructuredSemantics) { TEST_F(WireIteratorFixture, TraversalTerminatesAtFunctionReturn) { Value source; Value output; - auto module = + auto moduleOp = qco::QCOProgramBuilder::build(context.get(), [&](auto& builder) -> Value { source = builder.allocQubit(); output = builder.h(source); return output; }); - ASSERT_TRUE(module); + ASSERT_TRUE(moduleOp); qco::WireIterator it(source); ASSERT_EQ(it.operation(), source.getDefiningOp()); @@ -323,8 +323,8 @@ TEST_F(WireIteratorFixture, TraversalTerminatesAtFunctionReturn) { TEST_F(WireIteratorFixture, TraversalTerminatesAtUnknownCarrier) { OpBuilder builder(context.get()); const auto location = builder.getUnknownLoc(); - auto module = ModuleOp::create(location); - builder.setInsertionPointToStart(module.getBody()); + auto moduleOp = ModuleOp::create(location); + builder.setInsertionPointToStart(moduleOp.getBody()); auto function = func::FuncOp::create(builder, location, "main", builder.getFunctionType({}, {})); Block* body = function.addEntryBlock(); @@ -348,143 +348,49 @@ TEST_F(WireIteratorFixture, TraversalTerminatesAtUnknownCarrier) { EXPECT_EQ(backward, std::default_sentinel); } -TEST_F(WireIteratorFixture, CallMappingFollowsNestedReordering) { - auto module = parseModule(R"mlir( -func.func private @swap(%flag: i1, %a: !qco.qubit, %b: !qco.qubit) - -> (i1, !qco.qubit, !qco.qubit) { - return %flag, %b, %a : i1, !qco.qubit, !qco.qubit -} -func.func private @outer(%flag: i1, %a: !qco.qubit, %b: !qco.qubit) - -> (i1, !qco.qubit, !qco.qubit) { - %r:3 = func.call @swap(%flag, %a, %b) - : (i1, !qco.qubit, !qco.qubit) - -> (i1, !qco.qubit, !qco.qubit) - return %r#0, %r#1, %r#2 : i1, !qco.qubit, !qco.qubit -} -func.func @main() { - %flag = arith.constant true - %a = qco.alloc : !qco.qubit - %b = qco.alloc : !qco.qubit - %r:3 = func.call @outer(%flag, %a, %b) - : (i1, !qco.qubit, !qco.qubit) - -> (i1, !qco.qubit, !qco.qubit) - qco.sink %r#1 : !qco.qubit - qco.sink %r#2 : !qco.qubit - return -} -)mlir"); - ASSERT_TRUE(module); - auto main = module->lookupSymbol("main"); - auto call = findOp(main); - SmallVector allocs; - main.walk([&](qco::AllocOp op) { allocs.emplace_back(op.getResult()); }); - ASSERT_EQ(allocs.size(), 2U); - - qco::CallQubitMapping mapping; - auto mapped = mapping.getResultForOperand(call, call.getOperand(1)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_EQ(*mapped, call.getResult(2)); - mapped = mapping.getResultForOperand(call, call.getOperand(2)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_EQ(*mapped, call.getResult(1)); - - qco::WireIterator iterator(allocs[0]); +TEST_F(WireIteratorFixture, UnitaryCallContinuesWire) { + QCOProgramBuilder builder(context.get()); + builder.initialize(); + auto flip = builder.createUnitaryFunction( + "flip", TypeRange{QubitType::get(context.get())}, + [&](ValueRange arguments) { + return SmallVector{builder.x(arguments[0])}; + }); + Value input = builder.allocQubit(); + Value output = builder.call(flip, input).front(); + builder.sink(output); + auto moduleOp = builder.finalize(); + ASSERT_TRUE(moduleOp); + + WireIterator iterator(input); ++iterator; - EXPECT_EQ(iterator.qubit(), call.getResult(2)); + EXPECT_TRUE(isa(iterator.operation())); + EXPECT_EQ(iterator.qubit(), output); --iterator; - EXPECT_EQ(iterator.qubit(), allocs[0]); - - auto swap = module->lookupSymbol("swap"); - auto returnOp = cast(swap.getBody().front().getTerminator()); - returnOp->setOperands(swap.getArguments()); - mapping.invalidate(); - mapped = mapping.getResultForOperand(call, call.getOperand(1)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_EQ(*mapped, call.getResult(1)); + EXPECT_EQ(iterator.qubit(), input); } -TEST_F(WireIteratorFixture, CallMappingDistinguishesKeptAndCreatedQubits) { - auto module = parseModule(R"mlir( -func.func private @replace(%old: !qco.qubit) -> !qco.qubit { - qco.sink %old : !qco.qubit - %new = qco.alloc : !qco.qubit - return %new : !qco.qubit -} -func.func @main() { - %old = qco.alloc : !qco.qubit - %new = func.call @replace(%old) : (!qco.qubit) -> !qco.qubit - qco.sink %new : !qco.qubit - return -} -)mlir"); - ASSERT_TRUE(module); - auto main = module->lookupSymbol("main"); - auto call = findOp(main); - Value old = findOp(main).getResult(); - - qco::CallQubitMapping mapping; - auto mapped = mapping.getResultForOperand(call, old); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_FALSE(*mapped); - - qco::WireIterator consumed(old); - ++consumed; - ASSERT_EQ(consumed.operation(), call); - ++consumed; - EXPECT_EQ(consumed, std::default_sentinel); - - qco::WireIterator created(call.getResult(0)); - --created; - EXPECT_EQ(created, std::default_sentinel); -} +TEST_F(WireIteratorFixture, GenericCallIsWireBoundary) { + QCOProgramBuilder builder(context.get()); + builder.initialize(); + auto reset = builder.createFunction( + "reset", TypeRange{QubitType::get(context.get())}, + [&](ValueRange arguments) { + return SmallVector{builder.reset(arguments[0])}; + }); + Value input = builder.allocQubit(); + Value output = builder.call(reset, input).front(); + builder.sink(output); + auto moduleOp = builder.finalize(); + ASSERT_TRUE(moduleOp); -TEST_F(WireIteratorFixture, CallMappingFailsClosed) { - auto module = parseModule(R"mlir( -func.func private @external(!qco.qubit) -> !qco.qubit -func.func private @recursive(%q: !qco.qubit) -> !qco.qubit { - %r = func.call @recursive(%q) : (!qco.qubit) -> !qco.qubit - return %r : !qco.qubit -} -func.func private @unknown(%q: !qco.qubit) -> !qco.qubit { - %r = builtin.unrealized_conversion_cast %q : !qco.qubit to !qco.qubit - return %r : !qco.qubit -} -func.func @main() { - %a = qco.alloc : !qco.qubit - %x = func.call @external(%a) : (!qco.qubit) -> !qco.qubit - qco.sink %x : !qco.qubit - %b = qco.alloc : !qco.qubit - %y = func.call @recursive(%b) : (!qco.qubit) -> !qco.qubit - qco.sink %y : !qco.qubit - %c = qco.alloc : !qco.qubit - %z = func.call @unknown(%c) : (!qco.qubit) -> !qco.qubit - qco.sink %z : !qco.qubit - return -} -)mlir"); - ASSERT_TRUE(module); - auto main = module->lookupSymbol("main"); - func::CallOp external; - func::CallOp recursive; - func::CallOp unknown; - main.walk([&](func::CallOp call) { - if (call.getCallee() == "external") { - external = call; - } else if (call.getCallee() == "recursive") { - recursive = call; - } else { - unknown = call; - } - }); - ASSERT_TRUE(external); - ASSERT_TRUE(recursive); - ASSERT_TRUE(unknown); - - qco::CallQubitMapping mapping; - EXPECT_TRUE( - failed(mapping.getResultForOperand(external, external.getOperand(0)))); - EXPECT_TRUE( - failed(mapping.getResultForOperand(recursive, recursive.getOperand(0)))); - EXPECT_TRUE( - failed(mapping.getResultForOperand(unknown, unknown.getOperand(0)))); + WireIterator forward(input); + ++forward; + EXPECT_TRUE(isa(forward.operation())); + ++forward; + EXPECT_EQ(forward, std::default_sentinel); + + WireIterator backward(output); + --backward; + EXPECT_EQ(backward, std::default_sentinel); } diff --git a/mlir/unittests/Dialect/QTensor/Utils/test_tensoriterator.cpp b/mlir/unittests/Dialect/QTensor/Utils/test_tensoriterator.cpp index 6ba8311d41..b681f45b87 100644 --- a/mlir/unittests/Dialect/QTensor/Utils/test_tensoriterator.cpp +++ b/mlir/unittests/Dialect/QTensor/Utils/test_tensoriterator.cpp @@ -58,17 +58,6 @@ class TensorIteratorTest : public ::testing::Test { [[nodiscard]] OwningOpRef parseModule(StringRef source) const { return parseSourceString(source, context.get()); } - - [[nodiscard]] static func::CallOp findCall(Operation* root, - StringRef callee) { - func::CallOp found; - root->walk([&](func::CallOp call) { - if (call.getCallee() == callee) { - found = call; - } - }); - return found; - } }; } // namespace @@ -278,7 +267,7 @@ TEST_F(TensorIteratorTest, Traversal) { } TEST_F(TensorIteratorTest, CallResultStartsALifeChain) { - auto module = parseSourceString(R"mlir( + auto moduleOp = parseSourceString(R"mlir( func.func private @relabel(%t: tensor<2x!qco.qubit>) -> tensor<2x!qco.qubit> { return %t : tensor<2x!qco.qubit> } @@ -294,12 +283,12 @@ func.func @main() { return } )mlir", - context.get()); - ASSERT_TRUE(module); + context.get()); + ASSERT_TRUE(moduleOp); func::CallOp call; ExtractOp extract; - module->walk([&](Operation* op) { + moduleOp->walk([&](Operation* op) { if (auto c = dyn_cast(op)) { call = c; } @@ -354,12 +343,12 @@ TEST_F(TensorIteratorTest, TraversesMixedResultConditionals) { } )mlir"; - auto module = parseSourceString(source, context.get()); - ASSERT_TRUE(module); - ASSERT_TRUE(succeeded(verify(*module))); + auto moduleOp = parseSourceString(source, context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); qtensor::AllocOp alloc; - module->walk([&](qtensor::AllocOp candidate) { alloc = candidate; }); + moduleOp->walk([&](qtensor::AllocOp candidate) { alloc = candidate; }); ASSERT_TRUE(alloc); TensorIterator iterator(alloc.getResult()); @@ -462,107 +451,3 @@ TEST_F(TensorIteratorTest, TraversesWhileCarriedTensors) { ASSERT_EQ(swapped.operation(), tensor1.getDefiningOp()); ASSERT_EQ(swapped.tensor(), tensor1); } - -TEST_F(TensorIteratorTest, CallMappingFollowsNestedReordering) { - auto module = parseModule(R"mlir( -func.func private @swap( - %flag: i1, %a: tensor<2x!qco.qubit>, %b: tensor<2x!qco.qubit>) - -> (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) { - return %flag, %b, %a - : i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit> -} -func.func private @outer( - %flag: i1, %a: tensor<2x!qco.qubit>, %b: tensor<2x!qco.qubit>) - -> (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) { - %r:3 = func.call @swap(%flag, %a, %b) - : (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) - -> (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) - return %r#0, %r#1, %r#2 - : i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit> -} -func.func @main() { - %flag = arith.constant true - %c2 = arith.constant 2 : index - %a = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %b = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %r:3 = func.call @outer(%flag, %a, %b) - : (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) - -> (i1, tensor<2x!qco.qubit>, tensor<2x!qco.qubit>) - qtensor.dealloc %r#1 : tensor<2x!qco.qubit> - qtensor.dealloc %r#2 : tensor<2x!qco.qubit> - return -} -)mlir"); - ASSERT_TRUE(module); - auto main = module->lookupSymbol("main"); - auto call = findCall(main, "outer"); - ASSERT_TRUE(call); - - CallTensorMapping mapping; - auto mapped = mapping.getResultForOperand(call, call.getOperand(1)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_EQ(*mapped, call.getResult(2)); - mapped = mapping.getResultForOperand(call, call.getOperand(2)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_EQ(*mapped, call.getResult(1)); -} - -TEST_F(TensorIteratorTest, CallMappingReportsAKeptTensor) { - auto module = parseModule(R"mlir( -func.func private @consume(%t: tensor<2x!qco.qubit>) { - qtensor.dealloc %t : tensor<2x!qco.qubit> - return -} -func.func @main() { - %c2 = arith.constant 2 : index - %t = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - func.call @consume(%t) : (tensor<2x!qco.qubit>) -> () - return -} -)mlir"); - ASSERT_TRUE(module); - auto call = findCall(module->lookupSymbol("main"), "consume"); - ASSERT_TRUE(call); - - CallTensorMapping mapping; - auto mapped = mapping.getResultForOperand(call, call.getOperand(0)); - ASSERT_TRUE(succeeded(mapped)); - EXPECT_FALSE(*mapped); -} - -TEST_F(TensorIteratorTest, CallMappingFailsClosed) { - auto module = parseModule(R"mlir( -func.func private @external(tensor<2x!qco.qubit>) - -> tensor<2x!qco.qubit> -func.func private @recursive(%t: tensor<2x!qco.qubit>) - -> tensor<2x!qco.qubit> { - %r = func.call @recursive(%t) - : (tensor<2x!qco.qubit>) -> tensor<2x!qco.qubit> - return %r : tensor<2x!qco.qubit> -} -func.func @main() { - %c2 = arith.constant 2 : index - %a = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %x = func.call @external(%a) - : (tensor<2x!qco.qubit>) -> tensor<2x!qco.qubit> - qtensor.dealloc %x : tensor<2x!qco.qubit> - %b = qtensor.alloc(%c2) : tensor<2x!qco.qubit> - %y = func.call @recursive(%b) - : (tensor<2x!qco.qubit>) -> tensor<2x!qco.qubit> - qtensor.dealloc %y : tensor<2x!qco.qubit> - return -} -)mlir"); - ASSERT_TRUE(module); - auto main = module->lookupSymbol("main"); - auto external = findCall(main, "external"); - auto recursive = findCall(main, "recursive"); - ASSERT_TRUE(external); - ASSERT_TRUE(recursive); - - CallTensorMapping mapping; - EXPECT_TRUE( - failed(mapping.getResultForOperand(external, external.getOperand(0)))); - EXPECT_TRUE( - failed(mapping.getResultForOperand(recursive, recursive.getOperand(0)))); -} diff --git a/mlir/unittests/programs/qc_programs.cpp b/mlir/unittests/programs/qc_programs.cpp index 6ce30f017a..09f0574565 100644 --- a/mlir/unittests/programs/qc_programs.cpp +++ b/mlir/unittests/programs/qc_programs.cpp @@ -44,6 +44,26 @@ static Value measureAndReturn(QCProgramBuilder& b, ValueRange qubits) { Value emptyQC(QCProgramBuilder& b) { return b.intConstant(0); } +Value reusableUnitaryFunction(QCProgramBuilder& b) { + auto qubit = b.allocQubit(); + auto rotate = b.createUnitaryFunction( + "rotate", TypeRange{b.getF64Type(), qubit.getType()}, + [&](ValueRange arguments) { b.rx(arguments[0], arguments[1]); }); + b.call(rotate, {b.floatConstant(0.5), qubit}); + return b.measure(qubit); +} + +Value reusableResetFunction(QCProgramBuilder& b) { + auto qubit = b.allocQubit(); + auto reset = b.createFunction("reset", TypeRange{qubit.getType()}, + [&](ValueRange arguments) { + b.reset(arguments[0]); + return SmallVector{}; + }); + b.call(reset, qubit); + return b.measure(qubit); +} + Value allocQubit(QCProgramBuilder& b) { auto q = b.allocQubit(); return measureToRegister(b, q); diff --git a/mlir/unittests/programs/qc_programs.h b/mlir/unittests/programs/qc_programs.h index c6f6d2aeb3..03554b48ed 100644 --- a/mlir/unittests/programs/qc_programs.h +++ b/mlir/unittests/programs/qc_programs.h @@ -19,6 +19,12 @@ class QCProgramBuilder; /// Creates an empty QC Program. Value emptyQC(QCProgramBuilder& b); +/// Calls a reusable unitary rotation. +Value reusableUnitaryFunction(QCProgramBuilder& b); + +/// Calls a reusable reset function. +Value reusableResetFunction(QCProgramBuilder& b); + // --- Qubit Management ----------------------------------------------------- // /// Allocates a single qubit. diff --git a/mlir/unittests/programs/qco_programs.cpp b/mlir/unittests/programs/qco_programs.cpp index d2760ccf0c..fd16f1c28a 100644 --- a/mlir/unittests/programs/qco_programs.cpp +++ b/mlir/unittests/programs/qco_programs.cpp @@ -85,6 +85,27 @@ static Value measureAndReturn(QCOProgramBuilder& b, ValueRange qubits) { Value emptyQCO(QCOProgramBuilder& b) { return b.intConstant(0); } +Value reusableUnitaryFunction(QCOProgramBuilder& b) { + auto qubit = b.allocQubit(); + auto rotate = b.createUnitaryFunction( + "rotate", TypeRange{b.getF64Type(), qubit.getType()}, + [&](ValueRange arguments) { + return SmallVector{b.rx(arguments[0], arguments[1])}; + }); + qubit = b.call(rotate, {b.floatConstant(0.5), qubit}).back(); + return b.measure(qubit).second; +} + +Value reusableResetFunction(QCOProgramBuilder& b) { + auto qubit = b.allocQubit(); + auto reset = b.createFunction( + "reset", TypeRange{qubit.getType()}, [&](ValueRange arguments) { + return SmallVector{b.reset(arguments[0])}; + }); + qubit = b.call(reset, {qubit}).back(); + return b.measure(qubit).second; +} + Value allocQubit(QCOProgramBuilder& b) { auto q = b.allocQubit(); return measureToRegister(b, q); diff --git a/mlir/unittests/programs/qco_programs.h b/mlir/unittests/programs/qco_programs.h index 3ce3011864..55a18aaaf4 100644 --- a/mlir/unittests/programs/qco_programs.h +++ b/mlir/unittests/programs/qco_programs.h @@ -19,6 +19,12 @@ class QCOProgramBuilder; /// Creates an empty QCO program. Value emptyQCO(QCOProgramBuilder& builder); +/// Calls a reusable unitary rotation. +Value reusableUnitaryFunction(QCOProgramBuilder& b); + +/// Calls a reusable reset function. +Value reusableResetFunction(QCOProgramBuilder& b); + // --- Qubit Management ----------------------------------------------------- // /// Allocates a single qubit. diff --git a/noxfile.py b/noxfile.py index 83463dded8..29096c2be0 100755 --- a/noxfile.py +++ b/noxfile.py @@ -112,7 +112,6 @@ def cpp_lint(session: nox.Session) -> None: f"--files-changed-only={'false' if all_files else 'true'}", "--lines-changed-only=false", *(() if all_files else (f"--diff-base={diff_base}",)), - "--file-annotations=false", "--jobs=0", "--verbosity=info", env={"GITHUB_OUTPUT": str(output)},