MLIR 24.0.0git
FuncOps.cpp
Go to the documentation of this file.
1//===- FuncOps.cpp - Func Dialect Operations ------------------------------===//
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
10
15#include "mlir/IR/IRMapping.h"
16#include "mlir/IR/Matchers.h"
20#include "mlir/IR/Value.h"
23#include "llvm/ADT/APFloat.h"
24#include "llvm/ADT/MapVector.h"
25#include "llvm/ADT/STLExtras.h"
26#include "llvm/ADT/SmallVectorExtras.h"
27
28#include "mlir/Dialect/Func/IR/FuncOpsDialect.cpp.inc"
29
30using namespace mlir;
31using namespace mlir::func;
32
33//===----------------------------------------------------------------------===//
34// FuncDialect
35//===----------------------------------------------------------------------===//
36
37void FuncDialect::initialize() {
38 addOperations<
39#define GET_OP_LIST
40#include "mlir/Dialect/Func/IR/FuncOps.cpp.inc"
41 >();
42 declarePromisedInterface<ConvertToEmitCPatternInterface, FuncDialect>();
43 declarePromisedInterface<DialectInlinerInterface, FuncDialect>();
44 declarePromisedInterface<ConvertToLLVMPatternInterface, FuncDialect>();
45 declarePromisedInterfaces<bufferization::BufferizableOpInterface, CallOp,
46 FuncOp, ReturnOp>();
47}
48
49/// Materialize a single constant operation from a given attribute value with
50/// the desired resultant type.
51Operation *FuncDialect::materializeConstant(OpBuilder &builder, Attribute value,
52 Type type, Location loc) {
53 if (ConstantOp::isBuildableWith(value, type))
54 return ConstantOp::create(builder, loc, type,
55 llvm::cast<FlatSymbolRefAttr>(value));
56 return nullptr;
57}
58
59//===----------------------------------------------------------------------===//
60// CallOp
61//===----------------------------------------------------------------------===//
62
63LogicalResult CallOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
64 // Check that the callee attribute was specified.
65 auto fnAttr = (*this)->getAttrOfType<FlatSymbolRefAttr>("callee");
66 if (!fnAttr)
67 return emitOpError("requires a 'callee' symbol reference attribute");
68 FuncOp fn = symbolTable.lookupNearestSymbolFrom<FuncOp>(*this, fnAttr);
69 if (!fn)
70 return emitOpError() << "'" << fnAttr.getValue()
71 << "' does not reference a valid function";
72
73 // Verify that the operand and result types match the callee.
75}
76
77FunctionType CallOp::getCalleeType() {
78 return FunctionType::get(getContext(), getOperandTypes(), getResultTypes());
79}
80
81//===----------------------------------------------------------------------===//
82// CallIndirectOp
83//===----------------------------------------------------------------------===//
84
85/// Fold indirect calls that have a constant function as the callee operand.
86LogicalResult CallIndirectOp::canonicalize(CallIndirectOp indirectCall,
87 PatternRewriter &rewriter) {
88 // Check that the callee is a constant callee.
89 SymbolRefAttr calledFn;
90 if (!matchPattern(indirectCall.getCallee(), m_Constant(&calledFn)))
91 return failure();
92
93 // Replace with a direct call.
94 rewriter.replaceOpWithNewOp<CallOp>(indirectCall, calledFn,
95 indirectCall.getResultTypes(),
96 indirectCall.getArgOperands());
97 return success();
98}
99
100//===----------------------------------------------------------------------===//
101// ConstantOp
102//===----------------------------------------------------------------------===//
103
104LogicalResult ConstantOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
105 StringRef fnName = getValue();
106 Type type = getType();
107
108 // Try to find the referenced function.
109 auto fn = symbolTable.lookupNearestSymbolFrom<FuncOp>(
110 this->getOperation(), StringAttr::get(getContext(), fnName));
111 if (!fn)
112 return emitOpError() << "reference to undefined function '" << fnName
113 << "'";
114
115 // Check that the referenced function has the correct type.
116 if (fn.getFunctionType() != type)
117 return emitOpError("reference to function with mismatched type");
118
119 return success();
120}
121
122OpFoldResult ConstantOp::fold(FoldAdaptor adaptor) {
123 return getValueAttr();
124}
125
126void ConstantOp::getAsmResultNames(
127 function_ref<void(Value, StringRef)> setNameFn) {
128 setNameFn(getResult(), "f");
129}
130
131bool ConstantOp::isBuildableWith(Attribute value, Type type) {
132 return llvm::isa<FlatSymbolRefAttr>(value) && llvm::isa<FunctionType>(type);
133}
134
135//===----------------------------------------------------------------------===//
136// FuncOp
137//===----------------------------------------------------------------------===//
138
139FuncOp FuncOp::create(Location location, StringRef name, FunctionType type,
141 OpBuilder builder(location->getContext());
142 OperationState state(location, getOperationName());
143 FuncOp::build(builder, state, name, type, attrs);
144 return cast<FuncOp>(Operation::create(state));
145}
146FuncOp FuncOp::create(Location location, StringRef name, FunctionType type,
148 SmallVector<NamedAttribute, 8> attrRef(attrs);
149 return create(location, name, type, llvm::ArrayRef(attrRef));
150}
151FuncOp FuncOp::create(Location location, StringRef name, FunctionType type,
153 ArrayRef<DictionaryAttr> argAttrs) {
154 FuncOp func = create(location, name, type, attrs);
155 func.setAllArgAttrs(argAttrs);
156 return func;
157}
158
159void FuncOp::build(OpBuilder &builder, OperationState &state, StringRef name,
160 FunctionType type, ArrayRef<NamedAttribute> attrs,
161 ArrayRef<DictionaryAttr> argAttrs) {
163 builder.getStringAttr(name));
164 state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type));
165 state.attributes.append(attrs.begin(), attrs.end());
166 state.addRegion();
167
168 if (argAttrs.empty())
169 return;
170 assert(type.getNumInputs() == argAttrs.size());
172 builder, state, argAttrs, /*resultAttrs=*/{},
173 getArgAttrsAttrName(state.name), getResAttrsAttrName(state.name));
174}
175
176ParseResult FuncOp::parse(OpAsmParser &parser, OperationState &result) {
177 auto buildFuncType =
178 [](Builder &builder, ArrayRef<Type> argTypes, ArrayRef<Type> results,
180 std::string &) { return builder.getFunctionType(argTypes, results); };
181
183 parser, result, /*allowVariadic=*/false,
184 getFunctionTypeAttrName(result.name), buildFuncType,
185 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
186}
187
188void FuncOp::print(OpAsmPrinter &p) {
190 p, *this, /*isVariadic=*/false, getFunctionTypeAttrName(),
191 getArgAttrsAttrName(), getResAttrsAttrName());
192}
193
194/// Clone the internal blocks from this function into dest and all attributes
195/// from this function to dest.
196void FuncOp::cloneInto(FuncOp dest, IRMapping &mapper) {
197 // Add the attributes of this function to dest.
198 llvm::MapVector<StringAttr, Attribute> newAttrMap;
199 for (const auto &attr : dest->getAttrs())
200 newAttrMap.insert({attr.getName(), attr.getValue()});
201 for (const auto &attr : (*this)->getAttrs())
202 newAttrMap.insert({attr.getName(), attr.getValue()});
203
204 auto newAttrs = llvm::map_to_vector(
205 newAttrMap, [](std::pair<StringAttr, Attribute> attrPair) {
206 return NamedAttribute(attrPair.first, attrPair.second);
207 });
208 dest->setAttrs(DictionaryAttr::get(getContext(), newAttrs));
209
210 // Clone the body.
211 getBody().cloneInto(&dest.getBody(), mapper);
212}
213
214/// Create a deep copy of this function and all of its blocks, remapping
215/// any operands that use values outside of the function using the map that is
216/// provided (leaving them alone if no entry is present). Replaces references
217/// to cloned sub-values with the corresponding value that is copied, and adds
218/// those mappings to the mapper.
219FuncOp FuncOp::clone(IRMapping &mapper) {
220 // Create the new function.
221 FuncOp newFunc = cast<FuncOp>(getOperation()->cloneWithoutRegions());
222
223 // If the function has a body, then the user might be deleting arguments to
224 // the function by specifying them in the mapper. If so, we don't add the
225 // argument to the input type vector.
226 if (!isExternal()) {
227 FunctionType oldType = getFunctionType();
228
229 unsigned oldNumArgs = oldType.getNumInputs();
230 SmallVector<Type, 4> newInputs;
231 newInputs.reserve(oldNumArgs);
232 for (unsigned i = 0; i != oldNumArgs; ++i)
233 if (!mapper.contains(getArgument(i)))
234 newInputs.push_back(oldType.getInput(i));
235
236 /// If any of the arguments were dropped, update the type and drop any
237 /// necessary argument attributes.
238 if (newInputs.size() != oldNumArgs) {
239 newFunc.setType(FunctionType::get(oldType.getContext(), newInputs,
240 oldType.getResults()));
241
242 if (ArrayAttr argAttrs = getAllArgAttrs()) {
243 SmallVector<Attribute> newArgAttrs;
244 newArgAttrs.reserve(newInputs.size());
245 for (unsigned i = 0; i != oldNumArgs; ++i)
246 if (!mapper.contains(getArgument(i)))
247 newArgAttrs.push_back(argAttrs[i]);
248 newFunc.setAllArgAttrs(newArgAttrs);
249 }
250 }
251 }
252
253 /// Clone the current function into the new one and return it.
254 cloneInto(newFunc, mapper);
255 return newFunc;
256}
257FuncOp FuncOp::clone() {
258 IRMapping mapper;
259 return clone(mapper);
260}
261
262//===----------------------------------------------------------------------===//
263// ReturnOp
264//===----------------------------------------------------------------------===//
265
266LogicalResult FuncOp::verifyRegions() {
267 // External declarations have no body to check.
268 if (isDeclaration())
269 return success();
270 // Hoist the result types once; they are the same for every return site.
271 auto resultTypes = getFunctionType().getResults();
272 for (Block &block : getBody()) {
273 if (block.empty())
274 continue;
275 // Check func.return or other return-like terminators ops (e.g.
276 // llvm.return, test.return).
277 auto returnOp = dyn_cast<RegionBranchTerminatorOpInterface>(&block.back());
278 if (!returnOp)
279 continue;
280 auto operands =
281 returnOp.getMutableSuccessorOperands(RegionSuccessor(getOperation()));
282 if (operands.size() != resultTypes.size())
283 return returnOp->emitOpError("has ")
284 << operands.size() << " operands, but enclosing function (@"
285 << getName() << ") returns " << resultTypes.size();
286
287 for (auto [i, opType] : llvm::enumerate(llvm::zip(operands, resultTypes))) {
288 auto [operand, resTy] = opType;
289 if (operand.get().getType() != resTy)
290 return returnOp->emitError() << "type of return operand " << i << " ("
291 << operand.get().getType()
292 << ") doesn't match function result type ("
293 << resTy << ") in function @" << getName();
294 }
295 }
296
297 return success();
298}
299
300//===----------------------------------------------------------------------===//
301// TableGen'd op method definitions
302//===----------------------------------------------------------------------===//
303
304#define GET_OP_CLASSES
305#include "mlir/Dialect/Func/IR/FuncOps.cpp.inc"
return success()
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
ArrayAttr()
b getContext())
Attributes are known-constant values of operations.
Definition Attributes.h:25
MLIRContext * getContext() const
Return the context this attribute belongs to.
Block represents an ordered list of Operations.
Definition Block.h:33
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
FunctionType getFunctionType(TypeRange inputs, TypeRange results)
Definition Builders.cpp:84
StringAttr getStringAttr(const Twine &bytes)
Definition Builders.cpp:271
A symbol reference with a reference path containing a single element.
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
bool contains(T from) const
Checks to see if a mapping for 'from' exists.
Definition IRMapping.h:51
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
This class helps build Operations.
Definition Builders.h:210
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
iterator_range< dialect_attr_iterator > dialect_attr_range
Definition Operation.h:679
static Operation * create(Location location, OperationName name, TypeRange resultTypes, ValueRange operands, NamedAttrList &&attributes, PropertyRef properties, BlockRange successors, unsigned numRegions)
Create a new Operation with the specific fields.
Definition Operation.cpp:65
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class represents a successor of a region.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class represents a collection of SymbolTables.
virtual Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
static StringRef getSymbolAttrName()
Return the name of the attribute used for symbol names.
Definition SymbolTable.h:76
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
A named class for passing around the variadic flag.
LogicalResult verifyCallOpInterface(CallOpInterface call, TypeRange argumentTypes, TypeRange resultTypes)
Verify that the forwarded operands and results of call are in a 1:1 relationship with the given argum...
void addArgAndResultAttrs(Builder &builder, OperationState &result, ArrayRef< DictionaryAttr > argAttrs, ArrayRef< DictionaryAttr > resultAttrs, StringAttr argAttrsName, StringAttr resAttrsName)
Adds argument and result attributes, provided as argAttrs and resultAttrs arguments,...
void printFunctionOp(OpAsmPrinter &p, FunctionOpInterface op, bool isVariadic, StringRef typeAttrName, StringAttr argAttrsName, StringAttr resAttrsName)
Printer implementation for function-like operations.
ParseResult parseFunctionOp(OpAsmParser &parser, OperationState &result, bool allowVariadic, StringAttr typeAttrName, FuncTypeBuilder funcTypeBuilder, StringAttr argAttrsName, StringAttr resAttrsName)
Parser implementation for function-like operations.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
Operation * cloneWithoutRegions(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
This represents an operation in an abstracted form, suitable for use with the builder APIs.
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
Region * addRegion()
Create a region that should be attached to the operation.