17#include "llvm/Support/DebugLog.h"
20#define DEBUG_TYPE "convert-to-llvm"
23#define GEN_PASS_DEF_CONVERTTOLLVMPASS
24#include "mlir/Conversion/Passes.h.inc"
32class ConvertToLLVMPassInterface {
34 ConvertToLLVMPassInterface(MLIRContext *context,
35 ArrayRef<std::string> filterDialects,
36 bool allowPatternRollback =
true);
37 virtual ~ConvertToLLVMPassInterface() =
default;
40 static void getDependentDialects(DialectRegistry ®istry);
52 virtual LogicalResult transform(Operation *op,
53 AnalysisManager manager)
const = 0;
60 LogicalResult visitInterfaces(
61 llvm::function_ref<
void(ConvertToLLVMPatternInterface *)> visitor);
64 ArrayRef<std::string> filterDialects;
67 bool allowPatternRollback;
79 LoadDependentDialectExtension() : DialectExtensionBase({}) {}
81 void apply(MLIRContext *context,
82 MutableArrayRef<Dialect *> dialects)
const final {
83 LDBG() <<
"Convert to LLVM extension load";
85 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects)) {
86 LDBG() <<
"Convert to LLVM found dialect interface for "
87 << iface->getDialect()->getNamespace();
88 iface->loadDependentDialects(context);
93 std::unique_ptr<DialectExtensionBase>
clone() const final {
94 return std::make_unique<LoadDependentDialectExtension>(*
this);
104struct StaticConvertToLLVM :
public ConvertToLLVMPassInterface {
106 std::shared_ptr<const FrozenRewritePatternSet> patterns;
108 std::shared_ptr<const ConversionTarget> target;
110 std::shared_ptr<const LLVMTypeConverter> typeConverter;
111 using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
115 auto target = std::make_shared<ConversionTarget>(*context);
116 auto typeConverter = std::make_shared<LLVMTypeConverter>(context);
117 RewritePatternSet tempPatterns(context);
118 target->addLegalDialect<LLVM::LLVMDialect>();
120 if (
failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
121 iface->populateConvertToLLVMConversionPatterns(
122 *target, *typeConverter, tempPatterns);
126 std::make_unique<FrozenRewritePatternSet>(std::move(tempPatterns));
127 this->target = target;
128 this->typeConverter = typeConverter;
133 LogicalResult transform(Operation *op, AnalysisManager manager)
const final {
134 ConversionConfig config;
135 config.allowPatternRollback = allowPatternRollback;
136 if (
failed(applyPartialConversion(op, *target, *patterns, config)))
148struct DynamicConvertToLLVM :
public ConvertToLLVMPassInterface {
151 std::shared_ptr<const SmallVector<ConvertToLLVMPatternInterface *>>
153 using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
158 std::make_shared<SmallVector<ConvertToLLVMPatternInterface *>>();
160 if (
failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
161 interfaces->push_back(iface);
164 this->interfaces = interfaces;
169 LogicalResult transform(Operation *op, AnalysisManager manager)
const final {
170 RewritePatternSet patterns(context);
171 ConversionTarget
target(*context);
172 target.addLegalDialect<LLVM::LLVMDialect>();
174 const auto &dlAnalysis = manager.getAnalysis<DataLayoutAnalysis>();
175 const DataLayout &dl = dlAnalysis.getAtOrAbove(op);
176 LowerToLLVMOptions
options(context, dl);
177 LLVMTypeConverter typeConverter(context,
options, &dlAnalysis);
180 for (ConvertToLLVMPatternInterface *iface : *interfaces)
181 iface->populateConvertToLLVMConversionPatterns(
target, typeConverter,
189 ConversionConfig config;
190 config.allowPatternRollback = allowPatternRollback;
191 if (
failed(applyPartialConversion(op,
target, std::move(patterns), config)))
204class ConvertToLLVMPass
205 :
public impl::ConvertToLLVMPassBase<ConvertToLLVMPass> {
206 std::shared_ptr<const ConvertToLLVMPassInterface> impl;
209 using impl::ConvertToLLVMPassBase<ConvertToLLVMPass>::ConvertToLLVMPassBase;
210 void getDependentDialects(DialectRegistry ®istry)
const final {
211 ConvertToLLVMPassInterface::getDependentDialects(registry);
214 LogicalResult
initialize(MLIRContext *context)
final {
215 std::shared_ptr<ConvertToLLVMPassInterface> impl;
218 impl = std::make_shared<DynamicConvertToLLVM>(context, filterDialects,
219 allowPatternRollback);
221 impl = std::make_shared<StaticConvertToLLVM>(context, filterDialects,
222 allowPatternRollback);
223 if (
failed(impl->initialize()))
229 void runOnOperation() final {
230 if (
failed(impl->transform(getOperation(), getAnalysisManager())))
231 return signalPassFailure();
241ConvertToLLVMPassInterface::ConvertToLLVMPassInterface(
243 bool allowPatternRollback)
244 : context(context), filterDialects(filterDialects),
245 allowPatternRollback(allowPatternRollback) {}
247void ConvertToLLVMPassInterface::getDependentDialects(
249 registry.
insert<LLVM::LLVMDialect>();
253LogicalResult ConvertToLLVMPassInterface::visitInterfaces(
254 llvm::function_ref<
void(ConvertToLLVMPatternInterface *)> visitor) {
255 if (!filterDialects.empty()) {
259 for (StringRef dialectName : filterDialects) {
262 return emitError(UnknownLoc::get(context))
263 <<
"dialect not loaded: " << dialectName <<
"\n";
264 auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
266 return emitError(UnknownLoc::get(context))
267 <<
"dialect does not implement ConvertToLLVMPatternInterface: "
268 << dialectName <<
"\n";
278 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects))
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
static llvm::ManagedStatic< PassManagerOptions > options
#define MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(CLASS_NAME)
This class represents an opaque dialect extension.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
void addExtensions()
Add the given extensions to the registry.
MLIRContext is the top-level object for a collection of MLIR operations.
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
std::vector< Dialect * > getLoadedDialects()
Return information about all IR dialects loaded in the context.
Include the generated interface declarations.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
void registerConvertToLLVMDependentDialectLoading(DialectRegistry ®istry)
Register the extension that will load dependent dialects for LLVM conversion.
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
void populateOpConvertToLLVMConversionPatterns(Operation *op, ConversionTarget &target, LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
Helper function for populating LLVM conversion patterns.