MLIR 24.0.0git
ControlFlowToLLVM.cpp
Go to the documentation of this file.
1//===- ControlFlowToLLVM.cpp - ControlFlow to LLVM dialect 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//
9// This file implements a pass to convert MLIR standard and builtin dialects
10// into the LLVM IR dialect.
11//
12//===----------------------------------------------------------------------===//
13
15
24#include "mlir/IR/BuiltinOps.h"
26#include "mlir/Pass/Pass.h"
28
29namespace mlir {
30#define GEN_PASS_DEF_CONVERTCONTROLFLOWTOLLVMPASS
31#include "mlir/Conversion/Passes.h.inc"
32} // namespace mlir
33
34using namespace mlir;
35
36#define PASS_NAME "convert-cf-to-llvm"
37
38namespace {
39/// Lower `cf.assert`. The default lowering calls the `abort` function if the
40/// assertion is violated and has no effect otherwise. The failure message is
41/// ignored by the default lowering but should be propagated by any custom
42/// lowering.
43struct AssertOpLowering : public ConvertOpToLLVMPattern<cf::AssertOp> {
44 explicit AssertOpLowering(const LLVMTypeConverter &typeConverter,
45 bool abortOnFailedAssert = true,
46 SymbolTableCollection *symbolTables = nullptr)
47 : ConvertOpToLLVMPattern<cf::AssertOp>(typeConverter, /*benefit=*/1),
48 abortOnFailedAssert(abortOnFailedAssert), symbolTables(symbolTables) {}
49
50 LogicalResult
51 matchAndRewrite(cf::AssertOp op, OpAdaptor adaptor,
52 ConversionPatternRewriter &rewriter) const override {
53 auto loc = op.getLoc();
54 auto module = op->getParentOfType<ModuleOp>();
55
56 // Split block at `assert` operation.
57 Block *opBlock = rewriter.getInsertionBlock();
58 auto opPosition = rewriter.getInsertionPoint();
59 Block *continuationBlock = rewriter.splitBlock(opBlock, opPosition);
60
61 // Failed block: Generate IR to print the message and call `abort`.
62 Block *failureBlock = rewriter.createBlock(opBlock->getParent());
63 auto createResult = LLVM::createPrintStrCall(
64 rewriter, loc, module, "assert_msg", op.getMsg(), *getTypeConverter(),
65 /*addNewLine=*/false,
66 /*runtimeFunctionName=*/"puts", symbolTables);
67 if (createResult.failed())
68 return failure();
69
70 if (abortOnFailedAssert) {
71 // Insert the `abort` declaration if necessary.
72 auto abortFunc = module.lookupSymbol<LLVM::LLVMFuncOp>("abort");
73 if (!abortFunc) {
74 OpBuilder::InsertionGuard guard(rewriter);
75 rewriter.setInsertionPointToStart(module.getBody());
76 auto abortFuncTy = LLVM::LLVMFunctionType::get(getVoidType(), {});
77 abortFunc = LLVM::LLVMFuncOp::create(rewriter, rewriter.getUnknownLoc(),
78 "abort", abortFuncTy);
79 }
80 LLVM::CallOp::create(rewriter, loc, abortFunc, ValueRange());
81 LLVM::UnreachableOp::create(rewriter, loc);
82 } else {
83 LLVM::BrOp::create(rewriter, loc, ValueRange(), continuationBlock);
84 }
85
86 // Generate assertion test.
87 rewriter.setInsertionPointToEnd(opBlock);
88 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
89 op, adaptor.getArg(), continuationBlock, failureBlock);
90
91 return success();
92 }
93
94private:
95 /// If set to `false`, messages are printed but program execution continues.
96 /// This is useful for testing asserts.
97 bool abortOnFailedAssert = true;
98
99 SymbolTableCollection *symbolTables = nullptr;
100};
101
102/// Helper function for converting branch ops. This function converts the
103/// signature of the given block. If the new block signature is different from
104/// `expectedTypes`, returns "failure".
105static FailureOr<Block *> getConvertedBlock(ConversionPatternRewriter &rewriter,
106 const TypeConverter *converter,
107 Operation *branchOp, Block *block,
108 TypeRange expectedTypes) {
109 assert(converter && "expected non-null type converter");
110 assert(!block->isEntryBlock() && "entry blocks have no predecessors");
111
112 // There is nothing to do if the types already match.
113 if (block->getArgumentTypes() == expectedTypes)
114 return block;
115
116 // Compute the new block argument types and convert the block.
117 std::optional<TypeConverter::SignatureConversion> conversion =
118 converter->convertBlockSignature(block);
119 if (!conversion)
120 return rewriter.notifyMatchFailure(branchOp,
121 "could not compute block signature");
122 if (expectedTypes != conversion->getConvertedTypes())
123 return rewriter.notifyMatchFailure(
124 branchOp,
125 "mismatch between adaptor operand types and computed block signature");
126 return rewriter.applySignatureConversion(block, *conversion, converter);
127}
128
129/// Flatten the given value ranges into a single vector of values.
132 for (const ValueRange &vals : values)
133 llvm::append_range(result, vals);
134 return result;
135}
136
137/// Set attributes on an operation using its inherent/discardable split.
138static void setConvertedAttrs(Operation *op, DictionaryAttr attrs) {
139 SmallVector<NamedAttribute> discardableAttrs;
140 for (NamedAttribute attr : attrs) {
141 if (op->getInherentAttr(attr.getName()).has_value())
142 op->setInherentAttr(attr.getName(), attr.getValue());
143 else
144 discardableAttrs.push_back(attr);
145 }
146 op->setDiscardableAttrs(discardableAttrs);
147}
148
149/// Convert the destination block signature (if necessary) and lower the branch
150/// op to llvm.br.
151struct BranchOpLowering : public ConvertOpToLLVMPattern<cf::BranchOp> {
154
155 LogicalResult
156 matchAndRewrite(cf::BranchOp op, Adaptor adaptor,
157 ConversionPatternRewriter &rewriter) const override {
158 SmallVector<Value> flattenedAdaptor = flattenValues(adaptor.getOperands());
159 FailureOr<Block *> convertedBlock =
160 getConvertedBlock(rewriter, getTypeConverter(), op, op.getSuccessor(),
161 TypeRange(ValueRange(flattenedAdaptor)));
162 if (failed(convertedBlock))
163 return failure();
164 DictionaryAttr attrs = op->getDiscardableAttrDictionary();
165 auto loopAnnotation = op->getAttrOfType<LLVM::LoopAnnotationAttr>(
167 Operation *newOp = rewriter.replaceOpWithNewOp<LLVM::BrOp>(
168 op, flattenedAdaptor, loopAnnotation, *convertedBlock);
169 // TODO: We should not just forward all attributes like that. But there are
170 // existing Flang tests that depend on this behavior.
171 setConvertedAttrs(newOp, attrs);
172 return success();
173 }
174};
175
176/// Convert the destination block signatures (if necessary) and lower the
177/// branch op to llvm.cond_br.
178struct CondBranchOpLowering : public ConvertOpToLLVMPattern<cf::CondBranchOp> {
181
182 LogicalResult
183 matchAndRewrite(cf::CondBranchOp op, Adaptor adaptor,
184 ConversionPatternRewriter &rewriter) const override {
185 SmallVector<Value> flattenedAdaptorTrue =
186 flattenValues(adaptor.getTrueDestOperands());
187 SmallVector<Value> flattenedAdaptorFalse =
188 flattenValues(adaptor.getFalseDestOperands());
189 if (!llvm::hasSingleElement(adaptor.getCondition()))
190 return rewriter.notifyMatchFailure(op,
191 "expected single element condition");
192 FailureOr<Block *> convertedTrueBlock =
193 getConvertedBlock(rewriter, getTypeConverter(), op, op.getTrueDest(),
194 TypeRange(ValueRange(flattenedAdaptorTrue)));
195 if (failed(convertedTrueBlock))
196 return failure();
197 FailureOr<Block *> convertedFalseBlock =
198 getConvertedBlock(rewriter, getTypeConverter(), op, op.getFalseDest(),
199 TypeRange(ValueRange(flattenedAdaptorFalse)));
200 if (failed(convertedFalseBlock))
201 return failure();
202 DictionaryAttr attrs = op->getDiscardableAttrDictionary();
203 auto loopAnnotation = op->getAttrOfType<LLVM::LoopAnnotationAttr>(
205 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
206 op, llvm::getSingleElement(adaptor.getCondition()),
207 flattenedAdaptorTrue, flattenedAdaptorFalse, op.getBranchWeightsAttr(),
208 loopAnnotation, *convertedTrueBlock, *convertedFalseBlock);
209 // TODO: We should not just forward all attributes like that. But there are
210 // existing Flang tests that depend on this behavior.
211 setConvertedAttrs(newOp, attrs);
212 return success();
213 }
214};
215
216/// Convert the destination block signatures (if necessary) and lower the
217/// switch op to llvm.switch.
218struct SwitchOpLowering : public ConvertOpToLLVMPattern<cf::SwitchOp> {
220
221 LogicalResult
222 matchAndRewrite(cf::SwitchOp op, cf::SwitchOp::Adaptor adaptor,
223 ConversionPatternRewriter &rewriter) const override {
224 // Get or convert default block.
225 FailureOr<Block *> convertedDefaultBlock = getConvertedBlock(
226 rewriter, getTypeConverter(), op, op.getDefaultDestination(),
227 TypeRange(adaptor.getDefaultOperands()));
228 if (failed(convertedDefaultBlock))
229 return failure();
230
231 // Get or convert all case blocks.
232 SmallVector<Block *> caseDestinations;
233 SmallVector<ValueRange> caseOperands = adaptor.getCaseOperands();
234 for (auto it : llvm::enumerate(op.getCaseDestinations())) {
235 Block *b = it.value();
236 FailureOr<Block *> convertedBlock =
237 getConvertedBlock(rewriter, getTypeConverter(), op, b,
238 TypeRange(caseOperands[it.index()]));
239 if (failed(convertedBlock))
240 return failure();
241 caseDestinations.push_back(*convertedBlock);
242 }
243
244 rewriter.replaceOpWithNewOp<LLVM::SwitchOp>(
245 op, adaptor.getFlag(), *convertedDefaultBlock,
246 adaptor.getDefaultOperands(), adaptor.getCaseValuesAttr(),
247 caseDestinations, caseOperands);
248 return success();
249 }
250};
251
252} // namespace
253
255 const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
256 // clang-format off
257 patterns.add<
258 BranchOpLowering,
259 CondBranchOpLowering,
260 SwitchOpLowering>(converter);
261 // clang-format on
262}
263
265 const LLVMTypeConverter &converter, RewritePatternSet &patterns,
266 bool abortOnFailure, SymbolTableCollection *symbolTables) {
267 patterns.add<AssertOpLowering>(converter, abortOnFailure, symbolTables);
268}
269
270//===----------------------------------------------------------------------===//
271// Pass Definition
272//===----------------------------------------------------------------------===//
273
274namespace {
275/// A pass converting MLIR operations into the LLVM IR dialect.
276struct ConvertControlFlowToLLVM
277 : public impl::ConvertControlFlowToLLVMPassBase<ConvertControlFlowToLLVM> {
278
279 using Base::Base;
280
281 /// Run the dialect converter on the module.
282 void runOnOperation() override {
283 MLIRContext *ctx = &getContext();
285 // This pass lowers only CF dialect ops, but it also modifies block
286 // signatures inside other ops. These ops should be treated as legal. They
287 // are lowered by other passes.
288 target.markUnknownOpDynamicallyLegal([&](Operation *op) {
289 return op->getDialect() !=
290 ctx->getLoadedDialect<cf::ControlFlowDialect>();
291 });
292
293 const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
295 dataLayoutAnalysis.getAtOrAbove(getOperation()));
296 if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)
297 options.overrideIndexBitwidth(indexBitwidth);
298
299 LLVMTypeConverter converter(ctx, options, &dataLayoutAnalysis);
300 RewritePatternSet patterns(ctx);
303
304 if (failed(applyPartialConversion(getOperation(), target,
305 std::move(patterns))))
306 signalPassFailure();
307 }
308};
309} // namespace
310
311//===----------------------------------------------------------------------===//
312// ConvertToLLVMPatternInterface implementation
313//===----------------------------------------------------------------------===//
314
315namespace {
316/// Implement the interface to convert MemRef to LLVM.
317struct ControlFlowToLLVMDialectInterface
318 : public ConvertToLLVMPatternInterface {
319 ControlFlowToLLVMDialectInterface(Dialect *dialect)
320 : ConvertToLLVMPatternInterface(dialect) {}
321
322 void loadDependentDialects(MLIRContext *context) const final {
323 context->loadDialect<LLVM::LLVMDialect>();
324 }
325
326 /// Hook for derived dialect interface to provide conversion patterns
327 /// and mark dialect legal for the conversion target.
328 void populateConvertToLLVMConversionPatterns(
329 ConversionTarget &target, LLVMTypeConverter &typeConverter,
330 RewritePatternSet &patterns) const final {
332 patterns);
334 }
335};
336} // namespace
337
339 DialectRegistry &registry) {
340 registry.addExtension(+[](MLIRContext *ctx, cf::ControlFlowDialect *dialect) {
341 dialect->addInterfaces<ControlFlowToLLVMDialectInterface>();
342 });
343}
return success()
static SmallVector< Value > flattenValues(ArrayRef< ValueRange > values)
Flatten the given value ranges into a single vector of values.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
Block represents an ordered list of Operations.
Definition Block.h:34
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
Definition Block.cpp:154
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
Definition Block.cpp:27
bool isEntryBlock()
Return if this block is the entry block in the parent region.
Definition Block.cpp:36
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
typename SourceOp::template GenericAdaptor< ArrayRef< ValueRange > > OneToNOpAdaptor
Definition Pattern.h:236
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Derived class that automatically populates legalization information for different LLVM ops.
Conversion from types to the LLVM IR dialect.
Options to control the LLVM lowering.
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.
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
void setInherentAttr(StringAttr name, Attribute value)
Set an inherent attribute by name.
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Definition Operation.h:237
std::optional< Attribute > getInherentAttr(StringRef name)
Access an inherent attribute by name: returns an empty optional if there is no inherent attribute wit...
void setDiscardableAttrs(DictionaryAttr newAttrs)
Set the discardable attribute dictionary on this operation.
Definition Operation.h:575
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class represents a collection of SymbolTables.
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
constexpr llvm::StringLiteral getLoopAnnotationAttrName()
Canonical name used when attaching LoopAnnotationAttr as a discardable attribute on operations that d...
Definition LLVMAttrs.h:122
LogicalResult createPrintStrCall(OpBuilder &builder, Location loc, ModuleOp moduleOp, StringRef symbolName, StringRef string, const LLVMTypeConverter &typeConverter, bool addNewline=true, std::optional< StringRef > runtimeFunctionName={}, SymbolTableCollection *symbolTables=nullptr)
Generate IR that prints the given string to stdout.
void registerConvertControlFlowToLLVMInterface(DialectRegistry &registry)
void populateControlFlowToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
Collect the patterns to convert from the ControlFlow dialect to LLVM.
void populateAssertToLLVMConversionPattern(const LLVMTypeConverter &converter, RewritePatternSet &patterns, bool abortOnFailure=true, SymbolTableCollection *symbolTables=nullptr)
Populate the cf.assert to LLVM conversion pattern.
Include the generated interface declarations.
static constexpr unsigned kDeriveIndexBitwidthFromDataLayout
Value to pass as bitwidth for the index type when the converter is expected to derive the bitwidth fr...