MLIR 24.0.0git
ConvertToLLVMPass.cpp
Go to the documentation of this file.
1//===- ConvertToLLVMPass.cpp - MLIR LLVM Conversion -----------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
17#include "llvm/Support/DebugLog.h"
18#include <memory>
19
20#define DEBUG_TYPE "convert-to-llvm"
21
22namespace mlir {
23#define GEN_PASS_DEF_CONVERTTOLLVMPASS
24#include "mlir/Conversion/Passes.h.inc"
25} // namespace mlir
26
27using namespace mlir;
28
29namespace {
30/// Base class for creating the internal implementation of `convert-to-llvm`
31/// passes.
32class ConvertToLLVMPassInterface {
33public:
34 ConvertToLLVMPassInterface(MLIRContext *context,
35 ArrayRef<std::string> filterDialects,
36 bool allowPatternRollback = true);
37 virtual ~ConvertToLLVMPassInterface() = default;
38
39 /// Get the dependent dialects used by `convert-to-llvm`.
40 static void getDependentDialects(DialectRegistry &registry);
41
42 /// Initialize the internal state of the `convert-to-llvm` pass
43 /// implementation. This method is invoked by `ConvertToLLVMPass::initialize`.
44 /// This method returns whether the initialization process failed.
45 virtual LogicalResult initialize() = 0;
46
47 /// Transform `op` to LLVM with the conversions available in the pass. The
48 /// analysis manager can be used to query analyzes like `DataLayoutAnalysis`
49 /// to further configure the conversion process. This method is invoked by
50 /// `ConvertToLLVMPass::runOnOperation`. This method returns whether the
51 /// transformation process failed.
52 virtual LogicalResult transform(Operation *op,
53 AnalysisManager manager) const = 0;
54
55protected:
56 /// Visit the `ConvertToLLVMPatternInterface` dialect interfaces and call
57 /// `visitor` with each of the interfaces. If `filterDialects` is non-empty,
58 /// then `visitor` is invoked only with the dialects in the `filterDialects`
59 /// list.
60 LogicalResult visitInterfaces(
61 llvm::function_ref<void(ConvertToLLVMPatternInterface *)> visitor);
62 MLIRContext *context;
63 /// List of dialects names to use as filters.
64 ArrayRef<std::string> filterDialects;
65 /// An experimental flag to disallow pattern rollback. This is more efficient
66 /// but not supported by all lowering patterns.
67 bool allowPatternRollback;
68};
69
70/// This DialectExtension can be attached to the context, which will invoke the
71/// `apply()` method for every loaded dialect. If a dialect implements the
72/// `ConvertToLLVMPatternInterface` interface, we load dependent dialects
73/// through the interface. This extension is loaded in the context before
74/// starting a pass pipeline that involves dialect conversion to LLVM.
75class LoadDependentDialectExtension : public DialectExtensionBase {
76public:
77 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LoadDependentDialectExtension)
78
79 LoadDependentDialectExtension() : DialectExtensionBase(/*dialectNames=*/{}) {}
80
81 void apply(MLIRContext *context,
82 MutableArrayRef<Dialect *> dialects) const final {
83 LDBG() << "Convert to LLVM extension load";
84 for (auto *iface :
85 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects)) {
86 LDBG() << "Convert to LLVM found dialect interface for "
87 << iface->getDialect()->getNamespace();
88 iface->loadDependentDialects(context);
89 }
90 }
91
92 /// Return a copy of this extension.
93 std::unique_ptr<DialectExtensionBase> clone() const final {
94 return std::make_unique<LoadDependentDialectExtension>(*this);
95 }
96};
97
98//===----------------------------------------------------------------------===//
99// StaticConvertToLLVM
100//===----------------------------------------------------------------------===//
101
102/// Static implementation of the `convert-to-llvm` pass. This version only looks
103/// at dialect interfaces to configure the conversion process.
104struct StaticConvertToLLVM : public ConvertToLLVMPassInterface {
105 /// Pattern set with conversions to LLVM.
106 std::shared_ptr<const FrozenRewritePatternSet> patterns;
107 /// The conversion target.
108 std::shared_ptr<const ConversionTarget> target;
109 /// The LLVM type converter.
110 std::shared_ptr<const LLVMTypeConverter> typeConverter;
111 using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
112
113 /// Configure the conversion to LLVM at pass initialization.
114 LogicalResult initialize() final {
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>();
119 // Populate the patterns with the dialect interface.
120 if (failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
121 iface->populateConvertToLLVMConversionPatterns(
122 *target, *typeConverter, tempPatterns);
123 })))
124 return failure();
125 this->patterns =
126 std::make_unique<FrozenRewritePatternSet>(std::move(tempPatterns));
127 this->target = target;
128 this->typeConverter = typeConverter;
129 return success();
130 }
131
132 /// Apply the conversion driver.
133 LogicalResult transform(Operation *op, AnalysisManager manager) const final {
134 ConversionConfig config;
135 config.allowPatternRollback = allowPatternRollback;
136 if (failed(applyPartialConversion(op, *target, *patterns, config)))
137 return failure();
138 return success();
139 }
140};
141
142//===----------------------------------------------------------------------===//
143// DynamicConvertToLLVM
144//===----------------------------------------------------------------------===//
145
146/// Dynamic implementation of the `convert-to-llvm` pass. This version inspects
147/// the IR to configure the conversion to LLVM.
148struct DynamicConvertToLLVM : public ConvertToLLVMPassInterface {
149 /// A list of all the `ConvertToLLVMPatternInterface` dialect interfaces used
150 /// to partially configure the conversion process.
151 std::shared_ptr<const SmallVector<ConvertToLLVMPatternInterface *>>
152 interfaces;
153 using ConvertToLLVMPassInterface::ConvertToLLVMPassInterface;
154
155 /// Collect the dialect interfaces used to configure the conversion process.
156 LogicalResult initialize() final {
157 auto interfaces =
158 std::make_shared<SmallVector<ConvertToLLVMPatternInterface *>>();
159 // Collect the interfaces.
160 if (failed(visitInterfaces([&](ConvertToLLVMPatternInterface *iface) {
161 interfaces->push_back(iface);
162 })))
163 return failure();
164 this->interfaces = interfaces;
165 return success();
166 }
167
168 /// Configure the conversion process and apply the conversion driver.
169 LogicalResult transform(Operation *op, AnalysisManager manager) const final {
170 RewritePatternSet patterns(context);
171 ConversionTarget target(*context);
172 target.addLegalDialect<LLVM::LLVMDialect>();
173 // Get the data layout analysis.
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);
178
179 // Configure the conversion with dialect level interfaces.
180 for (ConvertToLLVMPatternInterface *iface : *interfaces)
181 iface->populateConvertToLLVMConversionPatterns(target, typeConverter,
182 patterns);
183
184 // Configure the conversion attribute interfaces.
186 patterns);
187
188 // Apply the conversion.
189 ConversionConfig config;
190 config.allowPatternRollback = allowPatternRollback;
191 if (failed(applyPartialConversion(op, target, std::move(patterns), config)))
192 return failure();
193 return success();
194 }
195};
196
197//===----------------------------------------------------------------------===//
198// ConvertToLLVMPass
199//===----------------------------------------------------------------------===//
200
201/// This is a generic pass to convert to LLVM, it uses the
202/// `ConvertToLLVMPatternInterface` dialect interface to delegate to dialects
203/// the injection of conversion patterns.
204class ConvertToLLVMPass
205 : public impl::ConvertToLLVMPassBase<ConvertToLLVMPass> {
206 std::shared_ptr<const ConvertToLLVMPassInterface> impl;
207
208public:
209 using impl::ConvertToLLVMPassBase<ConvertToLLVMPass>::ConvertToLLVMPassBase;
210 void getDependentDialects(DialectRegistry &registry) const final {
211 ConvertToLLVMPassInterface::getDependentDialects(registry);
212 }
213
214 LogicalResult initialize(MLIRContext *context) final {
215 std::shared_ptr<ConvertToLLVMPassInterface> impl;
216 // Choose the pass implementation.
217 if (useDynamic)
218 impl = std::make_shared<DynamicConvertToLLVM>(context, filterDialects,
219 allowPatternRollback);
220 else
221 impl = std::make_shared<StaticConvertToLLVM>(context, filterDialects,
222 allowPatternRollback);
223 if (failed(impl->initialize()))
224 return failure();
225 this->impl = impl;
226 return success();
227 }
228
229 void runOnOperation() final {
230 if (failed(impl->transform(getOperation(), getAnalysisManager())))
231 return signalPassFailure();
232 }
233};
234
235} // namespace
236
237//===----------------------------------------------------------------------===//
238// ConvertToLLVMPassInterface
239//===----------------------------------------------------------------------===//
240
241ConvertToLLVMPassInterface::ConvertToLLVMPassInterface(
242 MLIRContext *context, ArrayRef<std::string> filterDialects,
243 bool allowPatternRollback)
244 : context(context), filterDialects(filterDialects),
245 allowPatternRollback(allowPatternRollback) {}
246
247void ConvertToLLVMPassInterface::getDependentDialects(
248 DialectRegistry &registry) {
249 registry.insert<LLVM::LLVMDialect>();
250 registry.addExtensions<LoadDependentDialectExtension>();
251}
252
253LogicalResult ConvertToLLVMPassInterface::visitInterfaces(
254 llvm::function_ref<void(ConvertToLLVMPatternInterface *)> visitor) {
255 if (!filterDialects.empty()) {
256 // Test mode: Populate only patterns from the specified dialects. Produce
257 // an error if the dialect is not loaded or does not implement the
258 // interface.
259 for (StringRef dialectName : filterDialects) {
260 Dialect *dialect = context->getLoadedDialect(dialectName);
261 if (!dialect)
262 return emitError(UnknownLoc::get(context))
263 << "dialect not loaded: " << dialectName << "\n";
264 auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
265 if (!iface)
266 return emitError(UnknownLoc::get(context))
267 << "dialect does not implement ConvertToLLVMPatternInterface: "
268 << dialectName << "\n";
269 visitor(iface);
270 }
271 } else {
272 // Normal mode: Populate all patterns from all dialects that implement the
273 // interface.
274 // First time we encounter this dialect: if it implements the interface,
275 // let's populate patterns !
276 std::vector<Dialect *> dialects = context->getLoadedDialects();
277 for (auto *iface :
278 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects))
279 visitor(iface);
280 }
281 return success();
282}
283
284//===----------------------------------------------------------------------===//
285// API
286//===----------------------------------------------------------------------===//
287
289 DialectRegistry &registry) {
290 registry.addExtensions<LoadDependentDialectExtension>();
291}
return success()
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
static llvm::ManagedStatic< PassManagerOptions > options
#define MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(CLASS_NAME)
Definition TypeID.h:331
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.
Definition MLIRContext.h:63
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
void registerConvertToLLVMDependentDialectLoading(DialectRegistry &registry)
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.