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
2 changes: 1 addition & 1 deletion mlir/include/mlir/Conversion/CBitToMemRef/CBitToMemRef.td
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def ConvertCBitToMemRef : Pass<"convert-cbit-to-memref"> {
def ConvertCBitToMemRef : Pass<"convert-cbit-to-memref", "mlir::ModuleOp"> {
let summary = "Lower CBit registers to memrefs";

let description = [{
Expand Down
2 changes: 1 addition & 1 deletion mlir/include/mlir/Conversion/JeffToQCO/JeffToQCO.td
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def JeffToQCO : Pass<"jeff-to-qco"> {
def JeffToQCO : Pass<"jeff-to-qco", "mlir::ModuleOp"> {
let summary = "Convert `jeff` operations to QCO operations";

let description = [{
Expand Down
2 changes: 1 addition & 1 deletion mlir/include/mlir/Conversion/QCOToJeff/QCOToJeff.td
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def QCOToJeff : Pass<"qco-to-jeff"> {
def QCOToJeff : Pass<"qco-to-jeff", "mlir::ModuleOp"> {
let summary = "Convert QCO operations to `jeff` operations";

let description = [{
Expand Down
2 changes: 1 addition & 1 deletion mlir/include/mlir/Conversion/QCOToQC/QCOToQC.td
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def QCOToQC : Pass<"qco-to-qc"> {
def QCOToQC : Pass<"qco-to-qc", "mlir::ModuleOp"> {
let summary = "Convert QCO dialect to QC dialect.";

let description = [{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def QCToQIRAdaptive : Pass<"qc-to-qir-adaptive"> {
def QCToQIRAdaptive : Pass<"qc-to-qir-adaptive", "mlir::ModuleOp"> {
let summary = "Lower the QC dialect to the LLVM dialect compliant with the "
"QIR Adaptive Profile 2.1";

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

include "mlir/Pass/PassBase.td"

def QCToQIRBase : Pass<"qc-to-qir-base"> {
def QCToQIRBase : Pass<"qc-to-qir-base", "mlir::ModuleOp"> {
let summary = "Lower the QC dialect to the LLVM dialect compliant with the "
"QIR Base Profile 2.1";

Expand Down
2 changes: 1 addition & 1 deletion mlir/lib/Conversion/CBitToMemRef/CBitToMemRef.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,7 @@ struct ConvertCBitToMemRef final
protected:
void runOnOperation() override {
MLIRContext* context = &getContext();
auto* moduleOp = getOperation();
auto moduleOp = getOperation();
CBitTypeConverter typeConverter;
ConversionTarget target(*context);
RewritePatternSet patterns(context);
Expand Down
79 changes: 30 additions & 49 deletions mlir/lib/Conversion/JeffToQCO/JeffToQCO.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,6 @@
#include <jeff/IR/JeffOps.h>
#include <llvm/ADT/STLExtras.h>
#include <llvm/ADT/SmallVector.h>
#include <llvm/Support/ErrorHandling.h>
#include <mlir/Dialect/Arith/IR/Arith.h>
#include <mlir/Dialect/Func/IR/FuncOps.h>
#include <mlir/Dialect/Func/Transforms/FuncConversions.h>
Expand Down Expand Up @@ -204,55 +203,28 @@ static void createBarrierOp(jeff::CustomOp& op, jeff::CustomOpAdaptor& adaptor,
/**
* @brief Gets the name of the entry point from the module attributes
*/
static StringRef getEntryPointName(Operation* op) {
auto module = dyn_cast<ModuleOp>(op);
if (!module) {
llvm::reportFatalInternalError("Expected a module operation");
static FailureOr<StringRef> getEntryPointName(ModuleOp moduleOp) {
auto entryPointAttr = moduleOp->getAttrOfType<IntegerAttr>("jeff.entrypoint");
if (!entryPointAttr || !entryPointAttr.getType().isUnsignedInteger()) {
return moduleOp.emitError(
"requires an unsigned integer 'jeff.entrypoint' attribute");
}
auto entryPoint = entryPointAttr.getUInt();

auto entryPointAttr = module->getAttr("jeff.entrypoint");
if (!entryPointAttr) {
llvm::reportFatalInternalError(
"Module is missing 'jeff.entrypoint' attribute");
}
auto entryPoint = cast<IntegerAttr>(entryPointAttr).getUInt();

auto stringsAttr = module->getAttr("jeff.strings");
auto stringsAttr = moduleOp->getAttrOfType<ArrayAttr>("jeff.strings");
if (!stringsAttr) {
llvm::reportFatalInternalError(
"Module is missing 'jeff.strings' attribute");
return moduleOp.emitError("requires an array 'jeff.strings' attribute");
}
auto strings = cast<ArrayAttr>(stringsAttr);

if (entryPoint >= strings.size()) {
llvm::reportFatalInternalError("Entry point index is out of bounds");
if (entryPoint >= stringsAttr.size()) {
return moduleOp.emitError("'jeff.entrypoint' index is out of bounds");
}

return cast<StringAttr>(strings[entryPoint]).getValue();
}

/**
* @brief Cleans up the module after conversion
*
* @param op The module operation to clean up
* @return LogicalResult Success or failure of the cleanup
*/
static LogicalResult cleanUp(Operation* op) {
auto module = dyn_cast<ModuleOp>(op);
if (!module) {
return failure();
auto name = dyn_cast<StringAttr>(stringsAttr[entryPoint]);
if (!name) {
return moduleOp.emitError("'jeff.entrypoint' must index a string");
}

// Remove module attributes
module->removeAttr("jeff.entrypoint");
module->removeAttr("jeff.strings");
module->removeAttr("jeff.tool");
module->removeAttr("jeff.toolVersion");
module->removeAttr("jeff.version");
module->removeAttr("jeff.versionMinor");
module->removeAttr("jeff.versionPatch");

return success();
return name.getValue();
}

/**
Expand Down Expand Up @@ -1288,7 +1260,12 @@ struct JeffToQCO final : impl::JeffToQCOBase<JeffToQCO> {
protected:
void runOnOperation() override {
MLIRContext* context = &getContext();
auto* module = getOperation();
auto moduleOp = getOperation();
auto entryPointName = getEntryPointName(moduleOp);
if (failed(entryPointName)) {
signalPassFailure();
return;
}

ConversionTarget target(*context);
RewritePatternSet patterns(context);
Expand All @@ -1302,8 +1279,7 @@ struct JeffToQCO final : impl::JeffToQCOBase<JeffToQCO> {
tensor::TensorDialect, scf::SCFDialect>();

target.addDynamicallyLegalOp<func::FuncOp>([&](func::FuncOp op) {
return (op.getSymName() != getEntryPointName(module) ||
mqt::isEntryPoint(op)) &&
return (op.getSymName() != *entryPointName || mqt::isEntryPoint(op)) &&
typeConverter.isSignatureLegal(op.getFunctionType()) &&
typeConverter.isLegal(&op.getBody());
});
Expand Down Expand Up @@ -1343,14 +1319,19 @@ struct JeffToQCO final : impl::JeffToQCOBase<JeffToQCO> {
context);

// Apply the conversion
if (applyPartialConversion(module, target, std::move(patterns)).failed()) {
if (applyPartialConversion(moduleOp, target, std::move(patterns))
.failed()) {
signalPassFailure();
return;
}

if (cleanUp(module).failed()) {
signalPassFailure();
}
moduleOp->removeAttr("jeff.entrypoint");
moduleOp->removeAttr("jeff.strings");
moduleOp->removeAttr("jeff.tool");
moduleOp->removeAttr("jeff.toolVersion");
moduleOp->removeAttr("jeff.version");
moduleOp->removeAttr("jeff.versionMinor");
moduleOp->removeAttr("jeff.versionPatch");
}
};

Expand Down
41 changes: 18 additions & 23 deletions mlir/lib/Conversion/QCOToJeff/QCOToJeff.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -436,11 +436,11 @@ static void createPPROp(QCOOpType& op, ConversionPatternRewriter& rewriter,
}

/**
* @brief Updates all `jeff.yield` operations in @p module to use the latest
* @brief Updates all `jeff.yield` operations in @p moduleOp to use the latest
* classical-bit-register array values.
*/
static void patchCregYields(Operation* module, LoweringState& state) {
module->walk([&](jeff::YieldOp yieldOp) {
static void patchCregYields(ModuleOp moduleOp, LoweringState& state) {
moduleOp->walk([&](jeff::YieldOp yieldOp) {
auto* values = state.cbitState.getRegionValues(yieldOp->getParentRegion());
if (values == nullptr) {
return;
Expand All @@ -460,21 +460,16 @@ static void patchCregYields(Operation* module, LoweringState& state) {
/**
* @brief Cleans up the module after conversion
*
* @param op The module operation to clean up
* @param moduleOp The module operation to clean up
* @param state The lowering state
* @return LogicalResult Success or failure of the cleanup
*/
static LogicalResult cleanUp(Operation* op, LoweringState& state) {
static LogicalResult cleanUp(ModuleOp moduleOp, LoweringState& state) {
if (state.entryPointName.empty()) {
return failure();
}

auto module = dyn_cast<ModuleOp>(op);
if (!module) {
return failure();
}

for (auto funcOp : module.getOps<func::FuncOp>()) {
for (auto funcOp : moduleOp.getOps<func::FuncOp>()) {
state.strings.emplace_back(funcOp.getSymName());
}

Expand All @@ -488,26 +483,26 @@ static LogicalResult cleanUp(Operation* op, LoweringState& state) {
}
const auto entryPoint = static_cast<uint16_t>(distance);

// Set module attributes
OpBuilder builder(module.getContext());
OpBuilder builder(moduleOp.getContext());
auto uint16Type = builder.getIntegerType(16, false);

module->setAttr("jeff.entrypoint",
builder.getIntegerAttr(uint16Type, entryPoint));
moduleOp->setAttr("jeff.entrypoint",
builder.getIntegerAttr(uint16Type, entryPoint));

SmallVector<StringRef> stringRefs;
stringRefs.reserve(state.strings.size());
for (const auto& str : state.strings) {
stringRefs.emplace_back(str);
}
module->setAttr("jeff.strings", builder.getStrArrayAttr(stringRefs));
moduleOp->setAttr("jeff.strings", builder.getStrArrayAttr(stringRefs));

module->setAttr("jeff.tool", builder.getStringAttr("mqt-cc"));
module->setAttr("jeff.toolVersion", builder.getStringAttr(MQT_CORE_VERSION));
moduleOp->setAttr("jeff.tool", builder.getStringAttr("mqt-cc"));
moduleOp->setAttr("jeff.toolVersion",
builder.getStringAttr(MQT_CORE_VERSION));

module->setAttr("jeff.version", builder.getIntegerAttr(uint16Type, 0));
module->setAttr("jeff.versionMinor", builder.getIntegerAttr(uint16Type, 3));
module->setAttr("jeff.versionPatch", builder.getIntegerAttr(uint16Type, 0));
moduleOp->setAttr("jeff.version", builder.getIntegerAttr(uint16Type, 0));
moduleOp->setAttr("jeff.versionMinor", builder.getIntegerAttr(uint16Type, 3));
moduleOp->setAttr("jeff.versionPatch", builder.getIntegerAttr(uint16Type, 0));

return success();
}
Expand Down Expand Up @@ -1854,8 +1849,8 @@ struct QCOToJeff final : impl::QCOToJeffBase<QCOToJeff> {
protected:
void runOnOperation() override {
MLIRContext* context = &getContext();
auto* moduleOp = getOperation();
if (failed(mqt::normalizeGlobalPhases(cast<ModuleOp>(moduleOp)))) {
auto moduleOp = getOperation();
if (failed(mqt::normalizeGlobalPhases(moduleOp))) {
signalPassFailure();
return;
}
Expand Down
4 changes: 2 additions & 2 deletions mlir/lib/Conversion/QCOToQC/QCOToQC.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1174,7 +1174,7 @@ struct QCOToQC final : impl::QCOToQCBase<QCOToQC> {
protected:
void runOnOperation() override {
MLIRContext* context = &getContext();
auto* module = getOperation();
auto moduleOp = getOperation();

// Create state object to track the qubit addressing mode
LoweringState state;
Expand Down Expand Up @@ -1251,7 +1251,7 @@ struct QCOToQC final : impl::QCOToQCBase<QCOToQC> {
populateBranchOpInterfaceTypeConversionPattern(patterns, typeConverter);

// Apply the conversion
if (failed(applyPartialConversion(module, target, std::move(patterns)))) {
if (failed(applyPartialConversion(moduleOp, target, std::move(patterns)))) {
signalPassFailure();
}
}
Expand Down
4 changes: 2 additions & 2 deletions mlir/lib/Conversion/QCToQIR/QIRAdaptive/QCToQIRAdaptive.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -684,8 +684,8 @@ struct QCToQIRAdaptive final : impl::QCToQIRAdaptiveBase<QCToQIRAdaptive> {
*/
void runOnOperation() override {
MLIRContext* ctx = &getContext();
auto* moduleOp = getOperation();
if (failed(mqt::normalizeGlobalPhases(cast<ModuleOp>(moduleOp)))) {
auto moduleOp = getOperation();
if (failed(mqt::normalizeGlobalPhases(moduleOp))) {
signalPassFailure();
return;
}
Expand Down
18 changes: 10 additions & 8 deletions mlir/lib/Conversion/QCToQIR/QIRBase/QCToQIRBase.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -13,12 +13,12 @@
#include "mlir/Conversion/QCToQIR/QIRCommon/QIRCommon.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/MQT/Transforms/GlobalPhaseNormalization.h"
#include "mlir/Dialect/QC/IR/QCDialect.h"
#include "mlir/Dialect/QC/IR/QCOps.h"
#include "mlir/Dialect/QIR/Utils/QIRUtils.h"

#include <llvm/Support/ErrorHandling.h>
#include <mlir/Conversion/ArithToLLVM/ArithToLLVM.h>
#include <mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h>
#include <mlir/Conversion/FuncToLLVM/ConvertFuncToLLVM.h>
Expand Down Expand Up @@ -380,11 +380,6 @@ struct QCToQIRBase final : impl::QCToQIRBaseBase<QCToQIRBase> {
* @param main The main LLVM function to restructure
*/
static void ensureBlocks(LLVM::LLVMFuncOp& main, LoweringState& state) {
if (main.getBlocks().size() > 1) {
llvm::reportFatalInternalError(
"Modules with multiple blocks are not supported in the Base Profile");
}

// Get the existing block
auto* bodyBlock = &main.front();
OpBuilder builder(main.getBody());
Expand Down Expand Up @@ -461,8 +456,15 @@ struct QCToQIRBase final : impl::QCToQIRBaseBase<QCToQIRBase> {
*/
void runOnOperation() override {
MLIRContext* ctx = &getContext();
auto* moduleOp = getOperation();
if (failed(mqt::normalizeGlobalPhases(cast<ModuleOp>(moduleOp)))) {
auto moduleOp = getOperation();
auto entryPoint = mqt::getEntryPoint(moduleOp);
if (entryPoint && !entryPoint.getBody().hasOneBlock()) {
entryPoint.emitError(
"QIR Base Profile requires a single-block entry function");
signalPassFailure();
return;
}
if (failed(mqt::normalizeGlobalPhases(moduleOp))) {
signalPassFailure();
return;
}
Expand Down
Loading
Loading