30#define GEN_PASS_DEF_CONVERTCONTROLFLOWTOLLVMPASS
31#include "mlir/Conversion/Passes.h.inc"
36#define PASS_NAME "convert-cf-to-llvm"
45 bool abortOnFailedAssert =
true,
48 abortOnFailedAssert(abortOnFailedAssert), symbolTables(symbolTables) {}
51 matchAndRewrite(cf::AssertOp op, OpAdaptor adaptor,
52 ConversionPatternRewriter &rewriter)
const override {
53 auto loc = op.getLoc();
54 auto module = op->getParentOfType<ModuleOp>();
57 Block *opBlock = rewriter.getInsertionBlock();
58 auto opPosition = rewriter.getInsertionPoint();
59 Block *continuationBlock = rewriter.splitBlock(opBlock, opPosition);
64 rewriter, loc, module,
"assert_msg", op.getMsg(), *getTypeConverter(),
66 "puts", symbolTables);
67 if (createResult.failed())
70 if (abortOnFailedAssert) {
72 auto abortFunc =
module.lookupSymbol<LLVM::LLVMFuncOp>("abort");
75 rewriter.setInsertionPointToStart(module.getBody());
76 auto abortFuncTy = LLVM::LLVMFunctionType::get(getVoidType(), {});
77 abortFunc = LLVM::LLVMFuncOp::create(rewriter, rewriter.getUnknownLoc(),
78 "abort", abortFuncTy);
80 LLVM::CallOp::create(rewriter, loc, abortFunc,
ValueRange());
81 LLVM::UnreachableOp::create(rewriter, loc);
83 LLVM::BrOp::create(rewriter, loc,
ValueRange(), continuationBlock);
87 rewriter.setInsertionPointToEnd(opBlock);
88 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
89 op, adaptor.getArg(), continuationBlock, failureBlock);
97 bool abortOnFailedAssert =
true;
105static FailureOr<Block *> getConvertedBlock(ConversionPatternRewriter &rewriter,
109 assert(converter &&
"expected non-null type converter");
110 assert(!block->
isEntryBlock() &&
"entry blocks have no predecessors");
117 std::optional<TypeConverter::SignatureConversion> conversion =
118 converter->convertBlockSignature(block);
120 return rewriter.notifyMatchFailure(branchOp,
121 "could not compute block signature");
122 if (expectedTypes != conversion->getConvertedTypes())
123 return rewriter.notifyMatchFailure(
125 "mismatch between adaptor operand types and computed block signature");
126 return rewriter.applySignatureConversion(block, *conversion, converter);
133 llvm::append_range(
result, vals);
138static void setConvertedAttrs(
Operation *op, DictionaryAttr attrs) {
144 discardableAttrs.push_back(attr);
156 matchAndRewrite(cf::BranchOp op, Adaptor adaptor,
157 ConversionPatternRewriter &rewriter)
const override {
159 FailureOr<Block *> convertedBlock =
160 getConvertedBlock(rewriter, getTypeConverter(), op, op.getSuccessor(),
162 if (failed(convertedBlock))
164 DictionaryAttr attrs = op->getDiscardableAttrDictionary();
165 auto loopAnnotation = op->getAttrOfType<LLVM::LoopAnnotationAttr>(
167 Operation *newOp = rewriter.replaceOpWithNewOp<LLVM::BrOp>(
168 op, flattenedAdaptor, loopAnnotation, *convertedBlock);
171 setConvertedAttrs(newOp, attrs);
183 matchAndRewrite(cf::CondBranchOp op, Adaptor adaptor,
184 ConversionPatternRewriter &rewriter)
const override {
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(),
195 if (failed(convertedTrueBlock))
197 FailureOr<Block *> convertedFalseBlock =
198 getConvertedBlock(rewriter, getTypeConverter(), op, op.getFalseDest(),
200 if (failed(convertedFalseBlock))
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);
211 setConvertedAttrs(newOp, attrs);
222 matchAndRewrite(cf::SwitchOp op, cf::SwitchOp::Adaptor adaptor,
223 ConversionPatternRewriter &rewriter)
const override {
225 FailureOr<Block *> convertedDefaultBlock = getConvertedBlock(
226 rewriter, getTypeConverter(), op, op.getDefaultDestination(),
227 TypeRange(adaptor.getDefaultOperands()));
228 if (failed(convertedDefaultBlock))
234 for (
auto it : llvm::enumerate(op.getCaseDestinations())) {
236 FailureOr<Block *> convertedBlock =
237 getConvertedBlock(rewriter, getTypeConverter(), op,
b,
239 if (failed(convertedBlock))
241 caseDestinations.push_back(*convertedBlock);
244 rewriter.replaceOpWithNewOp<LLVM::SwitchOp>(
245 op, adaptor.getFlag(), *convertedDefaultBlock,
246 adaptor.getDefaultOperands(), adaptor.getCaseValuesAttr(),
247 caseDestinations, caseOperands);
259 CondBranchOpLowering,
260 SwitchOpLowering>(converter);
276struct ConvertControlFlowToLLVM
282 void runOnOperation()
override {
293 const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
295 dataLayoutAnalysis.getAtOrAbove(getOperation()));
297 options.overrideIndexBitwidth(indexBitwidth);
304 if (failed(applyPartialConversion(getOperation(),
target,
305 std::move(patterns))))
317struct ControlFlowToLLVMDialectInterface
318 :
public ConvertToLLVMPatternInterface {
319 ControlFlowToLLVMDialectInterface(Dialect *dialect)
320 : ConvertToLLVMPatternInterface(dialect) {}
322 void loadDependentDialects(MLIRContext *context)
const final {
323 context->loadDialect<LLVM::LLVMDialect>();
328 void populateConvertToLLVMConversionPatterns(
329 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
330 RewritePatternSet &patterns)
const final {
341 dialect->addInterfaces<ControlFlowToLLVMDialectInterface>();
static SmallVector< Value > flattenValues(ArrayRef< ValueRange > values)
Flatten the given value ranges into a single vector of values.
static llvm::ManagedStatic< PassManagerOptions > options
Block represents an ordered list of Operations.
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
bool isEntryBlock()
Return if this block is the entry block in the parent region.
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
typename SourceOp::template GenericAdaptor< ArrayRef< ValueRange > > OneToNOpAdaptor
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.
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
NamedAttribute represents a combination of a name and an Attribute value.
RAII guard to reset the insertion point of the builder when destroyed.
Operation is the basic unit of execution within MLIR.
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...
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.
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.
This class provides an abstraction over the different types of ranges over Values.
constexpr llvm::StringLiteral getLoopAnnotationAttrName()
Canonical name used when attaching LoopAnnotationAttr as a discardable attribute on operations that d...
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 ®istry)
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...