#include "mlir/Analysis/DataLayoutAnalysis.h"
#include "mlir/Conversion/ConvertToLLVM/ToLLVMInterface.h"
#include "mlir/Conversion/ConvertToLLVM/ToLLVMPass.h"
#include "mlir/Conversion/LLVMCommon/TypeConverter.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/Rewrite/FrozenRewritePatternSet.h"
#include "mlir/Transforms/DialectConversion.h"
#include "llvm/Support/DebugLog.h"
#include <memory>
#define DEBUG_TYPE "convert-to-llvm"
namespace mlir {
#define GEN_PASS_DEF_CONVERTTOLLVMPASS
#include "mlir/Conversion/Passes.h.inc"
}
using namespace mlir;
namespace {
class ConvertToLLVMPassInterface {
public:
ConvertToLLVMPassInterface(MLIRContext *context,
ArrayRef<std::string> filterDialects,
bool allowPatternRollback = true);
virtual ~ConvertToLLVMPassInterface() = default;
static void getDependentDialects(DialectRegistry ®istry);
virtual LogicalResult initialize() = 0;
virtual LogicalResult transform(Operation *op,
AnalysisManager manager) const = 0;
protected:
LogicalResult visitInterfaces(
llvm::function_ref<void(ConvertToLLVMPatternInterface *)> visitor);
MLIRContext *context;
ArrayRef<std::string> filterDialects;
bool allowPatternRollback;
};
class LoadDependentDialectExtension : public DialectExtensionBase {
public:
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LoadDependentDialectExtension)
LoadDependentDialectExtension() : DialectExtensionBase({}) {}
void apply(MLIRContext *context,
MutableArrayRef<Dialect *> dialects) const final {
LDBG() << "Convert to LLVM extension load";
for (Dialect *dialect : dialects) {
auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
if (!iface)
continue;
LDBG() << "Convert to LLVM found dialect interface for "
<< dialect->getNamespace();
iface->loadDependentDialects(context);
}
}
std::unique_ptr<DialectExtensionBase> clone() const final {
return std::make_unique<LoadDependentDialectExtension>(*this);
}
};
struct StaticConvertToLLVM : public ConvertToLLVMPassInterface {
std::shared_ptr<const FrozenRewritePatternSet> patterns;
std::shared_ptr<const ConversionTarget> target;
std::shared_ptr<const LLVMTypeConverter> typeConverter;
using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
LogicalResult initialize() final {
auto target = std::make_shared<ConversionTarget>(*context);
auto typeConverter = std::make_shared<LLVMTypeConverter>(context);
RewritePatternSet tempPatterns(context);
target->addLegalDialect<LLVM::LLVMDialect>();
if (failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
iface->populateConvertToLLVMConversionPatterns(
*target, *typeConverter, tempPatterns);
})))
return failure();
this->patterns =
std::make_unique<FrozenRewritePatternSet>(std::move(tempPatterns));
this->target = target;
this->typeConverter = typeConverter;
return success();
}
LogicalResult transform(Operation *op, AnalysisManager manager) const final {
ConversionConfig config;
config.allowPatternRollback = allowPatternRollback;
if (failed(applyPartialConversion(op, *target, *patterns, config)))
return failure();
return success();
}
};
struct DynamicConvertToLLVM : public ConvertToLLVMPassInterface {
std::shared_ptr<const SmallVector<ConvertToLLVMPatternInterface *>>
interfaces;
using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
LogicalResult initialize() final {
auto interfaces =
std::make_shared<SmallVector<ConvertToLLVMPatternInterface *>>();
if (failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
interfaces->push_back(iface);
})))
return failure();
this->interfaces = interfaces;
return success();
}
LogicalResult transform(Operation *op, AnalysisManager manager) const final {
RewritePatternSet patterns(context);
ConversionTarget target(*context);
target.addLegalDialect<LLVM::LLVMDialect>();
const auto &dlAnalysis = manager.getAnalysis<DataLayoutAnalysis>();
const DataLayout &dl = dlAnalysis.getAtOrAbove(op);
LowerToLLVMOptions options(context, dl);
LLVMTypeConverter typeConverter(context, options, &dlAnalysis);
for (ConvertToLLVMPatternInterface *iface : *interfaces)
iface->populateConvertToLLVMConversionPatterns(target, typeConverter,
patterns);
populateOpConvertToLLVMConversionPatterns(op, target, typeConverter,
patterns);
ConversionConfig config;
config.allowPatternRollback = allowPatternRollback;
if (failed(applyPartialConversion(op, target, std::move(patterns), config)))
return failure();
return success();
}
};
class ConvertToLLVMPass
: public impl::ConvertToLLVMPassBase<ConvertToLLVMPass> {
std::shared_ptr<const ConvertToLLVMPassInterface> impl;
public:
using impl::ConvertToLLVMPassBase<ConvertToLLVMPass>::ConvertToLLVMPassBase;
void getDependentDialects(DialectRegistry ®istry) const final {
ConvertToLLVMPassInterface::getDependentDialects(registry);
}
LogicalResult initialize(MLIRContext *context) final {
std::shared_ptr<ConvertToLLVMPassInterface> impl;
if (useDynamic)
impl = std::make_shared<DynamicConvertToLLVM>(context, filterDialects,
allowPatternRollback);
else
impl = std::make_shared<StaticConvertToLLVM>(context, filterDialects,
allowPatternRollback);
if (failed(impl->initialize()))
return failure();
this->impl = impl;
return success();
}
void runOnOperation() final {
if (failed(impl->transform(getOperation(), getAnalysisManager())))
return signalPassFailure();
}
};
}
ConvertToLLVMPassInterface::ConvertToLLVMPassInterface(
MLIRContext *context, ArrayRef<std::string> filterDialects,
bool allowPatternRollback)
: context(context), filterDialects(filterDialects),
allowPatternRollback(allowPatternRollback) {}
void ConvertToLLVMPassInterface::getDependentDialects(
DialectRegistry ®istry) {
registry.insert<LLVM::LLVMDialect>();
registry.addExtensions<LoadDependentDialectExtension>();
}
LogicalResult ConvertToLLVMPassInterface::visitInterfaces(
llvm::function_ref<void(ConvertToLLVMPatternInterface *)> visitor) {
if (!filterDialects.empty()) {
for (StringRef dialectName : filterDialects) {
Dialect *dialect = context->getLoadedDialect(dialectName);
if (!dialect)
return emitError(UnknownLoc::get(context))
<< "dialect not loaded: " << dialectName << "\n";
auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
if (!iface)
return emitError(UnknownLoc::get(context))
<< "dialect does not implement ConvertToLLVMPatternInterface: "
<< dialectName << "\n";
visitor(iface);
}
} else {
for (Dialect *dialect : context->getLoadedDialects()) {
auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
if (!iface)
continue;
visitor(iface);
}
}
return success();
}
void mlir::registerConvertToLLVMDependentDialectLoading(
DialectRegistry ®istry) {
registry.addExtensions<LoadDependentDialectExtension>();
}