diff --git a/CHANGELOG.md b/CHANGELOG.md index 0d0d690789..5ee30434b3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -48,8 +48,8 @@ releases may include breaking changes. [#1676], [#1706], [#1776], [#1836]) ([**@denialhaag**], [**@burgholzer**]) - ✨ Add a `place-and-route` pass for mapping circuits to architectures with restricted topologies ([#1537], [#1547], [#1568], [#1581], [#1583], [#1588], - [#1600], [#1664], [#1709], [#1716], [#1748], [#1805], [#1870], [#1904]) - ([**@MatthiasReumann**], [**@burgholzer**]) + [#1600], [#1664], [#1709], [#1716], [#1748], [#1805], [#1870], [#1904], + [#1911]) ([**@MatthiasReumann**], [**@burgholzer**]) - ✨ Add a pass for qubit reuse in quantum programs, as well as related auxiliary passes and patterns ([#1705], [#1755]) ([**@DRovara**]) - ✨ Add initial infrastructure for new QC and QCO MLIR dialects ([#1264], @@ -640,6 +640,7 @@ changelogs._ [#1914]: https://github.com/munich-quantum-toolkit/core/pull/1914 +[#1911]: https://github.com/munich-quantum-toolkit/core/pull/1911 [#1904]: https://github.com/munich-quantum-toolkit/core/pull/1904 [#1897]: https://github.com/munich-quantum-toolkit/core/pull/1897 [#1895]: https://github.com/munich-quantum-toolkit/core/pull/1895 diff --git a/mlir/include/mlir/Dialect/QCO/Utils/Drivers.h b/mlir/include/mlir/Dialect/QCO/Utils/Drivers.h index f359e3194c..6fe2730cbc 100644 --- a/mlir/include/mlir/Dialect/QCO/Utils/Drivers.h +++ b/mlir/include/mlir/Dialect/QCO/Utils/Drivers.h @@ -124,13 +124,13 @@ LogicalResult walkProgramGraph(MutableArrayRef wires, while (Traits::isActive(it)) { // For source-like (AllocOp, StaticOp, qtensor::ExtractOp), - // sink-like (SinkOp, YieldOp, qtensor::InsertOp, scf::YieldOp), and - // one-qubit non-unitary (ResetOp, MeasureOp) operations, simply advance - // the iterator. + // sink-like (SinkOp, YieldOp, qtensor::InsertOp, scf::YieldOp, + // scf::ConditionOp), and one-qubit non-unitary (ResetOp, MeasureOp) + // operations, simply advance the iterator. if (isa( - it.operation())) { + qtensor::ExtractOp, qtensor::InsertOp, scf::YieldOp, + scf::ConditionOp>(it.operation())) { std::ranges::advance(it, Traits::stride()); continue; } diff --git a/mlir/lib/Dialect/QCO/Transforms/Mapping/Mapping.cpp b/mlir/lib/Dialect/QCO/Transforms/Mapping/Mapping.cpp index 808131571a..b9ffe765b5 100644 --- a/mlir/lib/Dialect/QCO/Transforms/Mapping/Mapping.cpp +++ b/mlir/lib/Dialect/QCO/Transforms/Mapping/Mapping.cpp @@ -135,7 +135,7 @@ struct MappingPass : impl::MappingPassBase { /// Bidirectionally map a wire index to a program index. /// Overwrites existing mappings. - void map(const size_t index, const size_t prog) { + void insertOrUpdate(const size_t index, const size_t prog) { if (index >= indexToProgram_.size()) { indexToProgram_.resize(index + 1); } @@ -154,6 +154,9 @@ struct MappingPass : impl::MappingPassBase { std::swap(indexToProgram_[i0], indexToProgram_[i1]); } + /// Return the number of index-wire mappings. + [[nodiscard]] size_t size() const { return indexToProgram_.size(); } + private: /// Maps the i-th wire index to a program index. SmallVector indexToProgram_; @@ -437,6 +440,83 @@ struct MappingPass : impl::MappingPassBase { return newIfOp; } + /// Extend the arguments of an `scf::WhileOp` by adding a given range of + /// additional SSA values. Replaces the existing operation and returns the + /// newly created one. + static scf::WhileOp extend(scf::WhileOp whileOp, ValueRange addons, + IRRewriter& rewriter) { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(whileOp); + + Block* oldBefBlock = whileOp.getBeforeBody(); + Block* oldAftBlock = whileOp.getAfterBody(); + + const auto oldBefNumArgs = oldBefBlock->getNumArguments(); + const auto oldAftNumArgs = oldAftBlock->getNumArguments(); + + // Create a new while op at the same location as the old one with the + // additional arguments. + + SmallVector newInits(whileOp.getInits()); + newInits.append(addons.begin(), addons.end()); + + SmallVector newTypes(whileOp.getResultTypes()); + newTypes.append(addons.getTypes().begin(), addons.getTypes().end()); + + auto newWhileOp = + rewriter.create(whileOp.getLoc(), newTypes, newInits); + + const SmallVector locs(newTypes.size(), whileOp.getLoc()); + Block* newBefBlock = + rewriter.createBlock(&newWhileOp.getBefore(), {}, newTypes, locs); + Block* newAftBlock = + rewriter.createBlock(&newWhileOp.getAfter(), {}, newTypes, locs); + + rewriter.mergeBlocks(oldBefBlock, newBefBlock, + newBefBlock->getArguments().take_front(oldBefNumArgs)); + rewriter.mergeBlocks(oldAftBlock, newAftBlock, + newAftBlock->getArguments().take_front(oldAftNumArgs)); + + auto conditionOp = cast(newBefBlock->getTerminator()); + rewriter.setInsertionPoint(conditionOp); + + // Replace the old condition operation with one that includes the new + // "before" block arguments. + + SmallVector newConditionArgs(conditionOp.getArgs()); + llvm::append_range(newConditionArgs, + newBefBlock->getArguments().drop_front(oldBefNumArgs)); + + rewriter.create( + conditionOp.getLoc(), conditionOp.getCondition(), newConditionArgs); + rewriter.eraseOp(conditionOp); + + // Replace the old yield operation with one that includes the new "after" + // block arguments. + + auto yieldOp = cast(newAftBlock->getTerminator()); + rewriter.setInsertionPoint(yieldOp); + + SmallVector newYieldArgs(yieldOp.getResults()); + llvm::append_range(newYieldArgs, + newAftBlock->getArguments().drop_front(oldBefNumArgs)); + + rewriter.create(yieldOp.getLoc(), newYieldArgs); + rewriter.eraseOp(yieldOp); + + // Finally, replace the old while operation with the new one. + + rewriter.replaceOp( + whileOp, newWhileOp.getResults().take_front(whileOp.getNumResults())); + + for (const auto [before, after] : llvm::zip_equal( + addons, newWhileOp->getResults().take_back(addons.size()))) { + rewriter.replaceAllUsesExcept(before, after, newWhileOp); + } + + return newWhileOp; + } + /// Return the wires of a dynamic computation. /// The mapping pass currently assumes that /// - there are no `qco.alloc` operation @@ -475,7 +555,7 @@ struct MappingPass : impl::MappingPassBase { const auto index = wires.size(); wires.emplace_back(qubit); - infos.map(index, index); + infos.insertOrUpdate(index, index); continue; } @@ -533,7 +613,7 @@ struct MappingPass : impl::MappingPassBase { rewriter.eraseOp(op); wires.emplace_back(qubit); - infos.map(prog, prog); + infos.insertOrUpdate(prog, prog); }) .Case([&](auto op) { rewriter.setInsertionPointAfter(op); @@ -555,7 +635,7 @@ struct MappingPass : impl::MappingPassBase { const auto qubit = staticOps[hw].getQubit(); wires.emplace_back(qubit); - infos.map(prog, prog); + infos.insertOrUpdate(prog, prog); SinkOp::create(rewriter, body.getLoc(), qubit); } @@ -598,6 +678,28 @@ struct MappingPass : impl::MappingPassBase { DenseSet(newForOp.getRegionIterArgs().begin(), newForOp.getRegionIterArgs().end())); }) + .Case([&](scf::WhileOp whileOp) { + assert(qubits.size() == layout.nqubits()); + + llvm::for_each(whileOp.getInits(), + [&](Value v) { qubits.erase(v); }); + + auto newWhileOp = extend(whileOp, to_vector(qubits), rewriter); + for (const auto [init, result] : llvm::zip_equal( + newWhileOp.getInits(), newWhileOp.getResults())) { + qubits.insert(result); + qubits.erase(init); + } + + const auto beforeArgs = newWhileOp.getBeforeArguments(); + const auto afterArgs = newWhileOp.getAfterArguments(); + stack.emplace_back( + newWhileOp.getBefore(), + DenseSet(beforeArgs.begin(), beforeArgs.end())); + stack.emplace_back( + newWhileOp.getAfter(), + DenseSet(afterArgs.begin(), afterArgs.end())); + }) .Case([&](IfOp ifOp) { assert(qubits.size() == layout.nqubits()); @@ -1032,7 +1134,7 @@ struct MappingPass : impl::MappingPassBase { released.emplace_back(op); } }) - .template Case( + .template Case( [&](auto op) { stack.emplace_back(op, indices); }); } @@ -1046,6 +1148,27 @@ struct MappingPass : impl::MappingPassBase { return stack; } + /// Helper function to realign a terminator operation based on a permutation + /// of hardware indices. This constructs a value map from the given bundle and + /// reorders the terminator's operands according to the permutation vector. + template + static void realignTerminator(Operation* terminator, ArrayRef perm, + const RoutingBundle& bundle, + IRRewriter& rewriter, Args&&... extraArgs) { + // Map hardware indices to qubit values for the given bundle. + DenseMap m(bundle.wires.size()); + for (size_t i = 0; i < bundle.wires.size(); ++i) { + const auto prog = bundle.infos.lookupProgram(i); + const auto hw = bundle.layout.getHardwareIndex(prog); + m.try_emplace(hw, bundle.wires[i].qubit()); + } + + rewriter.setInsertionPoint(terminator); + rewriter.replaceOpWithNewOp( + terminator, std::forward(extraArgs)..., + to_vector(map_range(perm, [&](size_t hw) { return m.at(hw); }))); + } + /// Processes the recursive stack item by routing the nested operation and /// inserting epilogue SWAPs. template @@ -1054,171 +1177,208 @@ struct MappingPass : impl::MappingPassBase { RoutingBundle& parent, Statistics& stats, IRRewriter* rewriter = nullptr) { const auto& [op, indices] = item; - const LogicalResult res = - TypeSwitch(op) - .template Case([&](scf::ForOp forOp) { - RoutingBundle child{.layout = parent.layout}; - - // Construct child bundle. Particularly, find the iteration - // argument (block-argument, forward) or yielded result (backward) - // for each result qubit the selected (via indices) iterators - // point at. - - for (size_t i : indices) { - const auto prog = parent.infos.lookupProgram(i); - const auto res = cast(parent.wires[i].qubit()); - const auto arg = forOp.getTiedLoopRegionIterArg(res); - const auto index = child.wires.size(); - if constexpr (Direction == WireDirection::Forward) { - child.wires.emplace_back(arg); - } else { - const auto yield = forOp.getTiedLoopYieldedValue(arg)->get(); - child.wires.emplace_back(yield); - } - child.infos.map(index, prog); - } - - // Route the child body and prepare the wire iterators for - // epilogue SWAP insertion, i.e., point each iterator at the final - // qubit op (note: might be a measurement) before the yield: - // Because "route" moves each iterator to the default sentinel, - // decrement twice: sentinel → yield → unitary/block arg. + SmallVector permutation(indices.size()); + SmallVector children = + TypeSwitch>(op) + .template Case([&](auto) { + return SmallVector{ + RoutingBundle{.layout = parent.layout}}; + }) + .template Case([&](IfOp) { + return SmallVector{ + RoutingBundle{.layout = parent.layout}, + RoutingBundle{.layout = parent.layout}}; + }); - if (failed(route(child, stats, rewriter))) { - return failure(); + for (size_t i : indices) { + const auto prog = parent.infos.lookupProgram(i); + const auto hw = parent.layout.getHardwareIndex(prog); + const auto res = cast(parent.wires[i].qubit()); + const auto resNum = res.getResultNumber(); + + TypeSwitch(op) + .template Case([&](scf::ForOp forOp) { + children[0].infos.insertOrUpdate(children[0].infos.size(), prog); + children[0].wires.emplace_back([&] -> Value { + const auto arg = forOp.getTiedLoopRegionIterArg(res); + if constexpr (Direction == WireDirection::Forward) { + return arg; + } else { + return forOp.getTiedLoopYieldedValue(arg)->get(); } - - if constexpr (Mode == RoutingMode::Hot) { - for_each(child.wires, [](auto& it) { std::advance(it, -2); }); + }()); + }) + .template Case([&](scf::WhileOp whileOp) { + children[0].infos.insertOrUpdate(children[0].infos.size(), prog); + children[0].wires.emplace_back([&] -> Value { + auto condOp = cast( + whileOp.getBeforeBody()->getTerminator()); + const auto arg = whileOp.getBeforeArguments()[resNum]; + if constexpr (Direction == WireDirection::Forward) { + return arg; + } else { + return condOp.getArgs()[resNum]; } + }()); + }) + .template Case([&](IfOp ifOp) { + assert(ifOp.getNumRegions() == 2); + + OpOperand* const qubit = ifOp.getTiedQubit(res); + for (size_t i = 0; i < 2; ++i) { + Region& region = ifOp.getRegion(i); + + const auto arg = region.getRegionNumber() == 0 + ? ifOp.getTiedThenBlockArgument(qubit) + : ifOp.getTiedElseBlockArgument(qubit); + children[i].infos.insertOrUpdate(children[i].infos.size(), prog); + children[i].wires.emplace_back([&] -> Value { + if constexpr (Direction == WireDirection::Forward) { + return arg; + } else { + return region.getRegionNumber() == 0 + ? ifOp.getTiedThenYieldedValue(arg)->get() + : ifOp.getTiedElseYieldedValue(arg)->get(); + } + }()); + } + }); - // Find (insert) the epilogue SWAP sequence for (into) the child - // body using the "restore" strategy. Because this restores the - // parent's layout, we don't have to update the infos. - - const auto swaps = restore(child.layout, parent.layout); - insertSWAPs(swaps, child, stats, rewriter); - - // Sort topologically to fix any occurring SSA dominance errors. + permutation[resNum] = hw; + } - if constexpr (Mode == RoutingMode::Hot) { - sortTopologically(forOp.getBody()); - } + // Route each child branch and prepare the wire iterators for + // epilogue SWAP insertion, i.e., point each iterator at the final + // qubit op (note: might be a measurement) before the yield. + // TODO: Parallelize multiple children, if possible. - return success(); - }) - .template Case([&](IfOp ifOp) { - const std::array bodies{ - &ifOp.getThenRegion().getBlocks().front(), - &ifOp.getElseRegion().getBlocks().front(), - }; + for (auto& child : children) { + if (failed(route(child, stats, rewriter))) { + return failure(); + } - // Construct child bundles for each branch. Particularly, find the - // block-argument (forward) or yielded result (backward) for each - // result qubit the selected (via indices) iterators point at. + if constexpr (Mode == RoutingMode::Hot) { + for_each(child.wires, [](auto& it) { std::advance(it, -2); }); + } + } - SmallVector perm(indices.size()); - std::array children{RoutingBundle{.layout = parent.layout}, - RoutingBundle{.layout = parent.layout}}; + // Exception: The layout of the "after" region depends on the final layout + // of the before region. Thus, create / route the second child region / + // bundle here. - for (size_t i : indices) { - const auto prog = parent.infos.lookupProgram(i); - const auto res = cast(parent.wires[i].qubit()); - const auto index = children[0].wires.size(); + if (auto whileOp = dyn_cast(op)) { + children.emplace_back(RoutingBundle{.layout = children[0].layout}); + assert(children.size() == 2); - OpOperand* qubit = ifOp.getTiedQubit(res); - const std::array args{ifOp.getTiedThenBlockArgument(qubit), - ifOp.getTiedElseBlockArgument(qubit)}; + const auto rng = [&] -> ValueRange { + if constexpr (Direction == WireDirection::Forward) { + return whileOp.getAfterArguments(); + } + return cast(whileOp.getAfterBody()->getTerminator()) + .getResults(); + }(); + + for (const auto& [i, arg] : llvm::enumerate(rng)) { + const auto hw = permutation[i]; + const auto prog = children[0].layout.getProgramIndex(hw); + children[1].wires.emplace_back(arg); + children[1].infos.insertOrUpdate(i, prog); + } - perm[res.getResultNumber()] = - parent.layout.getHardwareIndex(prog); + if (failed(route(children[1], stats, rewriter))) { + return failure(); + } - if constexpr (Direction == WireDirection::Forward) { - for (size_t j = 0; j < children.size(); ++j) { - children[j].wires.emplace_back(args[j]); - children[j].infos.map(index, prog); - } - } else { - const std::array yields{ - ifOp.getTiedThenYieldedValue(args[0])->get(), - ifOp.getTiedElseYieldedValue(args[1])->get()}; - for (size_t j = 0; j < children.size(); ++j) { - children[j].wires.emplace_back(yields[j]); - children[j].infos.map(index, prog); - } - } - } + if constexpr (Mode == RoutingMode::Hot) { + for_each(children[1].wires, [](auto& it) { std::advance(it, -2); }); + } + } - // Route each child branch and prepare the wire iterators for - // epilogue SWAP insertion, i.e., point each iterator at the final - // qubit op (note: might be a measurement) before the yield. + const Layout exit = + TypeSwitch(op) + .Case([&](scf::ForOp) { + // Find (insert) the epilogue SWAP sequence for (into) the child + // region using the restore strategy. - for (auto& child : children) { - if (failed(route(child, stats, rewriter))) { - return failure(); - } + const auto swaps = restore(children[0].layout, parent.layout); + insertSWAPs(swaps, children[0], stats, rewriter); + return parent.layout; + }) + .template Case([&](scf::WhileOp) { + // Find (insert) the epilogue SWAP sequence for (into) the after + // region using the restore strategy. - if constexpr (Mode == RoutingMode::Hot) { - for_each(child.wires, [](auto& it) { std::advance(it, -2); }); - } - } + const auto swaps = restore(children[1].layout, parent.layout); + insertSWAPs(swaps, children[1], stats, rewriter); + // The scf::YieldOp is the terminator in the before region and + // thus determines the final output layout. + return children[0].layout; + }) + .template Case([&](IfOp) { // Find (insert) the epilogue SWAP sequence for (into) each child // branch using the "converge" strategy. const auto [convergedLayout, fst, snd] = converge(children[0].layout, children[1].layout); - insertSWAPs(fst, children[0], stats, rewriter); insertSWAPs(snd, children[1], stats, rewriter); + return convergedLayout; + }); - // Re-order the targets of a yield operation to match the input - // order of hardware indices and ensure qubits[i] = yield[i] and - // sort topologically to fix any occurring SSA dominance errors. - - if constexpr (Mode == RoutingMode::Hot) { - - for (const auto& [child, body] : - llvm::zip_equal(children, bodies)) { - - assert(all_of(child.wires, [&](auto& it) { - return isa(std::next(it).operation()); - })); - - DenseMap curr(child.wires.size()); - for (size_t i = 0; i < child.wires.size(); ++i) { - const auto prog = child.infos.lookupProgram(i); - const auto hw = child.layout.getHardwareIndex(prog); - curr.try_emplace(hw, child.wires[i].qubit()); - } - - auto yieldOp = cast(body->getTerminator()); - const SmallVector targets( - map_range(perm, [&](size_t hw) { return curr.at(hw); })); - rewriter->setInsertionPoint(yieldOp); - rewriter->replaceOpWithNewOp(yieldOp, targets); + if constexpr (Mode == RoutingMode::Hot) { + + // Realign terminator values to ensure that i-th input qubit and the i-th + // output qubit represent the equivalent hardware qubit. Note: Because + // we restore the layout at the end of the scf::ForOp, we don't require + // this procedure. + + TypeSwitch(op) + .template Case([&](scf::WhileOp whileOp) { + auto condOp = cast( + whileOp.getBeforeBody()->getTerminator()); + realignTerminator(condOp, permutation, + children[0], *rewriter, + condOp.getCondition()); + realignTerminator( + whileOp.getAfterBody()->getTerminator(), permutation, + children[1], *rewriter); + }) + .template Case([&](IfOp ifOp) { + realignTerminator( + ifOp.getThenRegion().front().getTerminator(), permutation, + children[0], *rewriter); + realignTerminator( + ifOp.getElseRegion().front().getTerminator(), permutation, + children[1], *rewriter); + }); + + // Sort topologically to fix any occurring SSA dominance errors. + + for (Region& region : op->getRegions()) { + assert(region.hasOneBlock()); + sortTopologically(®ion.front()); + } + } - sortTopologically(body); - } - } + // If the operation is a scf::ForOp, where the parent.layout = child.layout, + // we are done. Otherwise, propagate the final layout and index-to-program + // mapping to the parent. - // Propagate the correct layout and index-to-program mapping to - // the parent. + if (!isa(op)) { - WireInfos realigendInfos; - for (size_t i = 0; i < parent.wires.size(); ++i) { - const auto oldProg = parent.infos.lookupProgram(i); - const auto oldHw = parent.layout.getHardwareIndex(oldProg); - const auto newProg = convergedLayout.getProgramIndex(oldHw); - realigendInfos.map(i, newProg); - } + WireInfos realigendInfos; + for (size_t i = 0; i < parent.wires.size(); ++i) { + const auto oldProg = parent.infos.lookupProgram(i); + const auto oldHw = parent.layout.getHardwareIndex(oldProg); + const auto newProg = exit.getProgramIndex(oldHw); + realigendInfos.insertOrUpdate(i, newProg); + } - parent.layout = convergedLayout; - parent.infos = std::move(realigendInfos); - return success(); - }) - .Default([](Operation*) { return failure(); }); + parent.layout = exit; + parent.infos = std::move(realigendInfos); + } // Finally, move past the operation with nested regions by // incrementing the respective global wires. @@ -1227,7 +1387,7 @@ struct MappingPass : impl::MappingPassBase { std::advance(parent.wires[i], WireTraversalTraits::stride()); }); - return res; + return success(); } /// Iterates over a dynamically computed window of layers and uses A* search diff --git a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp index 8698895676..3b97c9ff01 100644 --- a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp +++ b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp @@ -28,7 +28,8 @@ namespace mlir::qco { bool WireIterator::isSinkLikeOperation(Operation* op) { - return isa(op); + return isa(op); } bool WireIterator::isSourceLikeOperation(Operation* op) { @@ -72,8 +73,16 @@ void WireIterator::forward() { }) .Case([&](MeasureOp op) { qubit_ = op.getQubitOut(); }) .Case([&](ResetOp op) { qubit_ = op.getQubitOut(); }) - .Case([&](auto op) { - qubit_ = op.getTiedLoopResult(&*(qubit_.use_begin())); + .Case([&](scf::ForOp op) { + qubit_ = op.getTiedLoopResult(qubit_.use_begin().getOperand()); + }) + .Case([&](scf::WhileOp op) { + // Because the scf::WhileOp doesn't implement "getLoopResults", we + // have to fallback to the following instead of using + // "getTiedLoopResult". + + OpOperand* operand = qubit_.use_begin().getOperand(); + qubit_ = op->getResult(operand->getOperandNumber()); }) .Case( [&](IfOp op) { qubit_ = op.getTiedResult(&(*qubit_.use_begin())); }) @@ -117,13 +126,25 @@ void WireIterator::backward() { [&](UnitaryOpInterface op) { qubit_ = op.getInputForOutput(qubit_); }) .Case([&](MeasureOp op) { qubit_ = op.getQubitIn(); }) .Case([&](ResetOp op) { qubit_ = op.getQubitIn(); }) - .Case([&](auto op) { + .Case([&](scf::ForOp op) { if (auto result = dyn_cast(qubit_)) { qubit_ = op.getTiedLoopInit(result)->get(); return; } llvm::reportFatalInternalError("expected result lookup"); }) + .Case([&](scf::WhileOp op) { + // Because the scf::WhileOp doesn't implement "getLoopResults", we + // have to fallback to the following instead of using + // "getTiedLoopInit". + + if (auto result = dyn_cast(qubit_)) { + qubit_ = op.getInits()[result.getResultNumber()]; + return; + } + + llvm::reportFatalInternalError("expected result lookup"); + }) .Case([&](IfOp op) { if (auto result = dyn_cast(qubit_)) { qubit_ = op.getTiedQubit(result)->get(); diff --git a/mlir/unittests/Dialect/QCO/Transforms/Mapping/test_mapping.cpp b/mlir/unittests/Dialect/QCO/Transforms/Mapping/test_mapping.cpp index 246d6aea50..89f8e45e84 100644 --- a/mlir/unittests/Dialect/QCO/Transforms/Mapping/test_mapping.cpp +++ b/mlir/unittests/Dialect/QCO/Transforms/Mapping/test_mapping.cpp @@ -18,6 +18,7 @@ #include #include +#include #include #include #include @@ -30,11 +31,11 @@ #include #include #include +#include #include #include #include -#include #include #include #include @@ -58,142 +59,115 @@ struct Device { static bool isExecutable(Region& body, DenseMap& m, const DenseSet>& couplingSet) { - for (Operation& rop : body.getOps()) { - const bool executable = - TypeSwitch(&rop) - .Case([&](StaticOp op) { - m.try_emplace(op.getQubit(), op.getIndex()); - return true; - }) - .Case([&](BarrierOp op) { - for (const auto [pred, succ] : - llvm::zip_equal(op.getInputQubits(), op.getOutputQubits())) { - m.try_emplace(succ, /*hw= */ m.at(pred)); - } - return true; - }) - .Case([&](UnitaryOpInterface& op) { - assert(op.getNumQubits() <= 2 && "expected two-qubit decomp."); - - if (op.getNumQubits() > 1) { - const auto hwA = m.at(op.getInputQubit(0)); - const auto hwB = m.at(op.getInputQubit(1)); - if (!couplingSet.contains(std::make_pair(hwA, hwB))) { - llvm::dbgs() << "The two-qubit gate (" << hwA << ", " << hwB - << ") is not executable: \n"; - op->dump(); - return false; - } - } - - for (const auto [pred, succ] : - llvm::zip_equal(op.getInputQubits(), op.getOutputQubits())) { - m.try_emplace(succ, /*hw= */ m.at(pred)); - } - - return true; - }) - .Case([&](scf::ForOp forOp) { - DenseMap bodyM; - for (const auto [init, arg] : llvm::zip_equal( - forOp.getInits(), forOp.getRegionIterArgs())) { - const auto hw = m.at(init); - bodyM.try_emplace(arg, hw); - } - - SmallVector initialHardwareOrder; - initialHardwareOrder.reserve(forOp.getInits().size()); - - for (OpOperand& operand : forOp.getInitsMutable()) { - const auto pred = operand.get(); - const auto succ = forOp.getTiedLoopResult(&operand); - const auto hw = m.at(pred); - - m.try_emplace(succ, hw); - initialHardwareOrder.emplace_back(hw); - } - - if (!isExecutable(forOp.getRegion(), bodyM, couplingSet)) { - return false; - } - - auto yield = cast(forOp.getBody()->getTerminator()); - - const SmallVector bodyHardwareOrder(llvm::map_range( - yield.getResults(), [&](auto v) { return bodyM.at(v); })); - - if (bodyHardwareOrder != initialHardwareOrder) { - llvm::dbgs() - << "The hardware indices of the yielded qubit values " - "must be in the same order as the scf::ForOp's " - "iteration qubit values!\n"; - return false; - } - - return true; - }) - .Case([&](qco::IfOp ifOp) { - std::array mappings{DenseMap{}, - DenseMap{}}; - - const std::array regions{&ifOp.getThenRegion(), - &ifOp.getElseRegion()}; - - for (size_t i = 0; i < 2; ++i) { - for (const auto [init, arg] : llvm::zip_equal( - ifOp.getQubits(), regions[i]->getArguments())) { - mappings[i].try_emplace(arg, /*hw = */ m.at(init)); - } - } - - SmallVector initialHardwareOrder; - initialHardwareOrder.reserve(ifOp.getQubits().size()); - - for (OpOperand& operand : ifOp.getQubitsMutable()) { - const auto pred = operand.get(); - const auto succ = ifOp.getTiedResult(&operand); - const auto hw = m.at(pred); - - m.try_emplace(succ, hw); - initialHardwareOrder.emplace_back(hw); - } - - for (const auto [body, mapping] : - llvm::zip_equal(regions, mappings)) { - if (!isExecutable(*body, mapping, couplingSet)) { - llvm::dbgs() - << "One of the qco::IfOp's branches is not executable!\n"; - return false; - } - - auto& block = body->getBlocks().front(); - auto yield = cast(block.getTerminator()); - - const SmallVector branchHardwareOrder(llvm::map_range( - yield.getTargets(), [&](auto v) { return mapping.at(v); })); - - if (branchHardwareOrder != initialHardwareOrder) { - llvm::dbgs() - << "The hardware indices of the yielded qubit values " - "must be in the same order as the qco::IfOp's input " - "qubit " - "values! This ensures that qco::IfOp's act like a " - "large, program-to-hardware mapping change, " - "unitary.\n"; - return false; - } - } - - return true; - }) - .Case([&](auto op) { - m.try_emplace(op.getQubitOut(), /*hw= */ m.at(op.getQubitIn())); - return true; - }) - .Default([](Operation*) { return true; }); - - if (!executable) { - return false; + for (Operation& op : body.getOps()) { + if (auto staticOp = dyn_cast(op)) { + m.try_emplace(staticOp.getQubit(), staticOp.getIndex()); + continue; + } + + if (auto unitaryOp = dyn_cast(op)) { + if (!isa(op) && unitaryOp.getNumQubits() > 1) { + assert(unitaryOp.getNumQubits() <= 2 && "expected two-qubit decomp."); + + const auto hwA = m.at(unitaryOp.getInputQubit(0)); + const auto hwB = m.at(unitaryOp.getInputQubit(1)); + if (!couplingSet.contains(std::make_pair(hwA, hwB))) { + llvm::dbgs() << "The two-qubit gate (" << hwA << ", " << hwB + << ") is not executable: \n"; + unitaryOp->dump(); + return false; + } + } + + for (const auto [pred, succ] : llvm::zip_equal( + unitaryOp.getInputQubits(), unitaryOp.getOutputQubits())) { + m.try_emplace(succ, m.at(pred)); + } + + continue; + } + + if (auto resetOp = dyn_cast(op)) { + m.try_emplace(resetOp.getQubitOut(), m.at(resetOp.getQubitIn())); + continue; + } + + if (auto measOp = dyn_cast(op)) { + m.try_emplace(measOp.getQubitOut(), m.at(measOp.getQubitIn())); + continue; + } + + if (!isa(op)) { + continue; + } + + for (Region& region : op.getRegions()) { + const ValueRange initArgs = + TypeSwitch(region.getParentOp()) + .Case([&](qco::IfOp ifOp) { return ifOp.getQubits(); }) + .Case( + [&](scf::WhileOp whileOp) { return whileOp.getInits(); }) + .Case( + [&](scf::ForOp forOp) { return forOp.getInits(); }) + .Default([](Operation*) -> ValueRange { return {}; }); + + const auto initialHardwareOrder = + to_vector(llvm::map_range(initArgs, [&](auto v) { return m.at(v); })); + + const auto qubitArgs = + llvm::make_filter_range(region.getArguments(), [](auto& arg) { + return isa(arg.getType()); + }); + + DenseMap localM; + for (const auto [i, arg] : llvm::enumerate(qubitArgs)) { + localM.try_emplace(arg, initialHardwareOrder[i]); + } + + if (!isExecutable(region, localM, couplingSet)) { + return false; + } + + Operation* terminator = region.front().getTerminator(); + const ValueRange finalOrderArgs = + TypeSwitch(region.getParentOp()) + .Case([&](qco::IfOp) { + return cast(terminator).getTargets(); + }) + .Case([&](auto) { + // Choose between "before" and "after" terminator. + return region.getRegionNumber() == 0 + ? cast(terminator).getArgs() + : cast(terminator).getResults(); + }) + .Case([&](scf::ForOp) { + return cast(terminator).getResults(); + }) + .Default([](Operation*) -> ValueRange { return {}; }); + + const auto finalOrder = to_vector(llvm::map_range( + finalOrderArgs, [&](auto v) { return localM.at(v); })); + + if (finalOrder != initialHardwareOrder) { + llvm::dbgs() + << "The hardware indices of the yielded terminator qubit values " + "must be in the same order as parent's op input qubit values!\n"; + return false; + } + } + + for (OpResult res : op.getResults()) { + const Value init = TypeSwitch(&op) + .Case([&](scf::WhileOp whileOp) { + return whileOp.getInits()[res.getResultNumber()]; + }) + .Case([&](scf::ForOp forOp) { + return forOp.getTiedLoopInit(res)->get(); + }) + .Case([&](qco::IfOp ifOp) { + return ifOp.getTiedQubit(res)->get(); + }); + m.try_emplace(res, m.at(init)); } } @@ -217,8 +191,8 @@ static Device getNineQubitSquareGrid() { {5, 8}, {8, 5}, {6, 7}, {7, 6}, {7, 8}, {8, 7}}}; } -/// Creates an N-qubit GHZ state, where N = `qubits.size()` using straight-line -/// programming. +/// Creates an N-qubit GHZ state, where N = `qubits.size()` using +/// straight-line programming. static void flatGHZ(QCOProgramBuilder& builder, SmallVector& qubits) { qubits[0] = builder.h(qubits[0]); for (size_t i = 1; i < qubits.size(); ++i) { @@ -579,7 +553,6 @@ TEST_P(MappingPassTest, MapParallelLoops) { for (int64_t i = 0; i < size; ++i) { std::tie(qubits[i], bits[i]) = builder.measure(qubits[i]); - qubits[i] = builder.h(qubits[i]); } for (int64_t i = 0; i < size; ++i) { @@ -722,5 +695,68 @@ TEST_P(MappingPassTest, MapBranchingGHZ) { EXPECT_TRUE(isExecutable(entry, device.couplingSet)); } +TEST_P(MappingPassTest, MapDoUntil) { + const auto& device = GetParam(); + const auto size = 4; + + QCOProgramBuilder builder(context.get()); + builder.initialize(); + + Value tensor = builder.qtensorAlloc(size); + SmallVector qubits(size); + + for (int64_t i = 0; i < size; ++i) { + std::tie(tensor, qubits[i]) = builder.qtensorExtract(tensor, i); + } + + qubits = builder.scfWhile( + qubits, + [&](ValueRange args) { + SmallVector beforeArgs(args); + SmallVector beforeBits(args); + + flatGHZ(builder, beforeArgs); + + beforeArgs = builder.barrier(beforeArgs); + + for (int64_t i = 0; i < size; ++i) { + std::tie(beforeArgs[i], beforeBits[i]) = + builder.measure(beforeArgs[i]); + } + + for (int64_t i = 0; i < size - 1; ++i) { + beforeBits[i + 1] = + arith::AndIOp::create(builder, beforeBits[i], beforeBits[i + 1]) + .getResult(); + } + + builder.scfCondition(beforeBits[size - 1], beforeArgs); + return beforeArgs; + }, + [&](ValueRange args) { + SmallVector afterArgs(args); + flatGHZ(builder, afterArgs); + return afterArgs; + }); + + flatGHZ(builder, qubits); + + qubits = builder.barrier(qubits); + + for (int64_t i = 0; i < size; ++i) { + tensor = builder.qtensorInsert(qubits[i], tensor, i); + } + + builder.qtensorDealloc(tensor); + + auto m = builder.finalize(); + auto res = + runPass(m.get(), device.couplingSet, MappingPassOptions{.ntrials = 1}); + auto entry = getEntryPoint(m.get()); + + ASSERT_TRUE(res.succeeded()); + EXPECT_TRUE(isExecutable(entry, device.couplingSet)); +} + INSTANTIATE_TEST_SUITE_P(NineQubitSquareGrid, MappingPassTest, testing::Values(getNineQubitSquareGrid()));