Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion mlir/include/mlir/Dialect/QTensor/IR/QTensorDialect.td
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,10 @@ def QTensorDialect : Dialect {
In addition, alloc/dealloc operations are added to the dialect to support the bulk allocation and deallocation of qubit tensors with linear types.
}];

let dependentDialects = ["::mlir::qco::QCODialect"];
let dependentDialects = ["::mlir::qco::QCODialect",
"::mlir::arith::ArithDialect"];

let hasCanonicalizer = 1;

let cppNamespace = "::mlir::qtensor";
}
Expand Down
240 changes: 2 additions & 238 deletions mlir/lib/Dialect/QCO/IR/SCF/IfOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,17 +10,14 @@

#include "mlir/Dialect/QCO/IR/QCOOps.h"
#include "mlir/Dialect/QCO/QCOUtils.h"
#include "mlir/Dialect/QTensor/IR/QTensorOps.h"

#include <llvm/ADT/BitVector.h>
#include <llvm/ADT/DenseMap.h>
#include <llvm/ADT/STLExtras.h>
#include <llvm/ADT/STLFunctionalExtras.h>
#include <llvm/ADT/Sequence.h>
#include <llvm/ADT/SmallPtrSet.h>
#include <llvm/ADT/SmallVector.h>
#include <mlir/Dialect/Arith/IR/Arith.h>
#include <mlir/Dialect/Utils/StaticValueUtils.h>
#include <mlir/IR/Attributes.h>
#include <mlir/IR/Builders.h>
#include <mlir/IR/BuiltinTypes.h>
Expand All @@ -37,7 +34,6 @@
#include <cassert>
#include <cstddef>
#include <cstdint>
#include <optional>

using namespace mlir;
using namespace mlir::qco;
Expand Down Expand Up @@ -334,244 +330,12 @@ struct RemoveUnusedClassicalResults : public OpRewritePattern<IfOp> {
}
};

struct QTensorAccess {
qtensor::ExtractOp extract;
qtensor::InsertOp insert;
};

struct BranchQTensorAccesses {
DenseMap<int64_t, QTensorAccess> accesses;
SmallVector<Operation*> qTensorOperations;
};

} // namespace

/// Analyze a QTensor's complete lifetime in one branch.
///
/// Supported branches extract distinct constant-index qubits, perform
/// QTensor-independent computation, reinsert one qubit at every extracted
/// index, and yield the resulting QTensor. Dynamic indices, repeated accesses,
/// and partial updates do not match.
static std::optional<BranchQTensorAccesses>
analyzeQTensorBranch(Block* block, size_t qTensorArgumentIndex,
size_t qTensorYieldIndex) {
BranchQTensorAccesses result;
Value currentQTensor = block->getArgument(qTensorArgumentIndex);
bool reachedInsertPhase = false;

while (true) {
assert(currentQTensor.hasOneUse() && "expected linear typing");
Operation* user = *currentQTensor.getUsers().begin();
if (user->getBlock() != block) {
return std::nullopt;
}

if (auto extract = dyn_cast<qtensor::ExtractOp>(user)) {
auto index = getConstantIntValue(extract.getIndex());
if (reachedInsertPhase || !index ||
!result.accesses
.try_emplace(*index, QTensorAccess{.extract = extract})
.second) {
return std::nullopt;
}
result.qTensorOperations.push_back(user);
currentQTensor = extract.getOutTensor();
continue;
}

if (auto insert = dyn_cast<qtensor::InsertOp>(user)) {
reachedInsertPhase = true;
auto index = getConstantIntValue(insert.getIndex());
if (!index) {
return std::nullopt;
}
auto access = result.accesses.find(*index);
if (access == result.accesses.end() || access->second.insert) {
return std::nullopt;
}
access->second.insert = insert;
result.qTensorOperations.push_back(user);
currentQTensor = insert.getResult();
continue;
}

auto yield = dyn_cast<YieldOp>(user);
if (!yield || user != block->getTerminator() ||
qTensorYieldIndex >= yield.getTargets().size() ||
yield.getTargets()[qTensorYieldIndex] != currentQTensor ||
llvm::any_of(result.accesses, [](const auto& access) {
return !access.second.insert;
})) {
return std::nullopt;
}
return result;
}
}

/// Move a branch while replacing QTensor accesses with scalar qubits.
static void moveScalarizedQTensorBranch(IfOp oldIf, Block* oldBlock,
Block* newBlock,
size_t qTensorArgumentIndex,
BranchQTensorAccesses& accesses,
ArrayRef<int64_t> indices,
PatternRewriter& rewriter) {
auto oldYield = cast<YieldOp>(oldBlock->getTerminator());
auto scalarArguments = newBlock->getArguments().take_back(indices.size());
auto carriedArguments = newBlock->getArguments().drop_back(indices.size());

SmallVector<Value> argumentReplacements;
argumentReplacements.reserve(oldBlock->getNumArguments());
size_t carriedIndex = 0;
for (size_t oldIndex : llvm::seq(oldBlock->getNumArguments())) {
argumentReplacements.push_back(oldIndex == qTensorArgumentIndex
? oldIf.getQubits()[qTensorArgumentIndex]
: carriedArguments[carriedIndex++]);
}
assert(carriedIndex == carriedArguments.size());
rewriter.mergeBlocks(oldBlock, newBlock, argumentReplacements);

SmallVector<Value> scalarYields;
scalarYields.reserve(indices.size());
for (auto [indexPosition, index] : llvm::enumerate(indices)) {
auto access = accesses.accesses.find(index);
if (access == accesses.accesses.end()) {
scalarYields.push_back(scalarArguments[indexPosition]);
} else {
rewriter.replaceAllUsesWith(access->second.extract.getResult(),
scalarArguments[indexPosition]);
scalarYields.push_back(access->second.insert.getScalar());
}
}

auto oldTargets = oldYield.getTargets();
size_t classicalResultCount = oldIf.getClassicalResults().size();
SmallVector<Value> newYieldValues;
newYieldValues.reserve(oldTargets.size() - 1 + scalarYields.size());
llvm::append_range(newYieldValues,
oldTargets.take_front(classicalResultCount));
for (auto [oldIndex, value] :
llvm::enumerate(oldTargets.drop_front(classicalResultCount))) {
if (oldIndex != qTensorArgumentIndex) {
newYieldValues.push_back(value);
}
}
llvm::append_range(newYieldValues, scalarYields);

rewriter.setInsertionPoint(oldYield);
rewriter.replaceOpWithNewOp<YieldOp>(oldYield, newYieldValues);

for (Operation* operation : llvm::reverse(accesses.qTensorOperations)) {
rewriter.eraseOp(operation);
}
}

namespace {

/// Replace constant-index QTensor updates in an if with scalar threading.
///
/// A QTensor carried through an if hides its qubits from target mapping. This
/// pattern extracts the union of constant indices accessed by either branch,
/// threads those qubits through both branches, and reinserts the results.
/// Untouched elements remain in the QTensor outside the if.
struct ScalarizeQTensorInputs final : OpRewritePattern<IfOp> {
using OpRewritePattern::OpRewritePattern;

LogicalResult matchAndRewrite(IfOp op,
PatternRewriter& rewriter) const override {
size_t classicalResultCount = op.getClassicalResults().size();
auto oldQubits = op.getQubits();

for (auto [qTensorIndex, qTensor] : llvm::enumerate(oldQubits)) {
auto qTensorType = dyn_cast<RankedTensorType>(qTensor.getType());
if (!qTensorType || !qTensorType.hasStaticShape()) {
continue;
}

auto thenAccesses = analyzeQTensorBranch(
op.thenBlock(), qTensorIndex, classicalResultCount + qTensorIndex);
auto elseAccesses = analyzeQTensorBranch(
op.elseBlock(), qTensorIndex, classicalResultCount + qTensorIndex);
if (!thenAccesses || !elseAccesses) {
continue;
}

SmallVector<int64_t> accessedIndices(thenAccesses->accesses.keys());
llvm::append_range(accessedIndices, elseAccesses->accesses.keys());
llvm::sort(accessedIndices);
accessedIndices.erase(llvm::unique(accessedIndices),
accessedIndices.end());
ArrayRef<int64_t> indices(accessedIndices);

rewriter.setInsertionPoint(op);
SmallVector<Value> indexValues;
SmallVector<Value> scalarInputs;
indexValues.reserve(indices.size());
scalarInputs.reserve(indices.size());
Value qTensorWithoutScalars = qTensor;
for (int64_t index : indices) {
auto indexValue =
arith::ConstantIndexOp::create(rewriter, op.getLoc(), index);
auto extract = qtensor::ExtractOp::create(rewriter, op.getLoc(),
qTensorWithoutScalars,
indexValue.getResult());
indexValues.push_back(indexValue.getResult());
scalarInputs.push_back(extract.getResult());
qTensorWithoutScalars = extract.getOutTensor();
}

SmallVector<Value> newQubits(oldQubits);
newQubits.erase(newQubits.begin() + qTensorIndex);
llvm::append_range(newQubits, scalarInputs);

auto newIf = IfOp::create(
rewriter, op.getLoc(), op.getClassicalResults().getTypes(),
ValueRange(newQubits).getTypes(), op.getCondition(), newQubits);
newIf->setDiscardableAttrs(op->getDiscardableAttrDictionary());

SmallVector<Location> locations(newQubits.size(), op.getLoc());
Block* oldThenBlock = op.thenBlock();
Block* oldElseBlock = op.elseBlock();
Block* newThenBlock =
rewriter.createBlock(&newIf.getThenRegion(), {},
ValueRange(newQubits).getTypes(), locations);
Block* newElseBlock =
rewriter.createBlock(&newIf.getElseRegion(), {},
ValueRange(newQubits).getTypes(), locations);
moveScalarizedQTensorBranch(op, oldThenBlock, newThenBlock, qTensorIndex,
*thenAccesses, indices, rewriter);
moveScalarizedQTensorBranch(op, oldElseBlock, newElseBlock, qTensorIndex,
*elseAccesses, indices, rewriter);

rewriter.setInsertionPointAfter(newIf);
Value updatedQTensor = qTensorWithoutScalars;
auto scalarResults = newIf.getLinearResults().take_back(indices.size());
for (auto [scalar, indexValue] :
llvm::zip_equal(scalarResults, indexValues)) {
updatedQTensor =
qtensor::InsertOp::create(rewriter, op.getLoc(), scalar,
updatedQTensor, indexValue)
.getResult();
}

SmallVector<Value> replacements(
newIf.getLinearResults().drop_back(indices.size()));
replacements.insert(replacements.begin() + qTensorIndex, updatedQTensor);
replacements.insert(replacements.begin(),
newIf.getClassicalResults().begin(),
newIf.getClassicalResults().end());
rewriter.replaceOp(op, replacements);
return success();
}
return failure();
}
};
} // namespace

void IfOp::getCanonicalizationPatterns(RewritePatternSet& results,
MLIRContext* context) {
results
.add<RemoveStaticCondition, ConditionPropagation, ForwardClassicalResults,
RemoveUnusedClassicalResults, ScalarizeQTensorInputs>(context);
results.add<RemoveStaticCondition, ConditionPropagation,
ForwardClassicalResults, RemoveUnusedClassicalResults>(context);
}

LogicalResult IfOp::verify() {
Expand Down
2 changes: 1 addition & 1 deletion mlir/lib/Dialect/QIR/Execution/Runtime/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -21,5 +21,5 @@ if(NOT TARGET ${TARGET_NAME})
target_link_libraries(
${TARGET_NAME}
PUBLIC MLIRQCODDAdapter MQT::CoreDD
PRIVATE MLIRQTensorDialect MQT::ProjectWarnings MQT::ProjectOptions)
PRIVATE MQT::ProjectWarnings MQT::ProjectOptions)
endif()
2 changes: 1 addition & 1 deletion mlir/lib/Dialect/QTensor/IR/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ file(GLOB_RECURSE OPERATIONS "${CMAKE_CURRENT_SOURCE_DIR}/Operations/*.cpp")

add_mlir_dialect_library(
MLIRQTensorDialect
QTensorCanonicalization.cpp
QTensorOps.cpp
${OPERATIONS}
ADDITIONAL_HEADER_DIRS
Expand All @@ -18,7 +19,6 @@ add_mlir_dialect_library(
MLIRQTensorOpsIncGen
LINK_LIBS
PRIVATE
MLIRQTensorUtils
MLIRIR
MLIRDialectUtils
MLIRArithDialect
Expand Down
Loading
Loading