21#include <mlir/IR/PatternMatch.h>
22#include <mlir/IR/SymbolTable.h>
23#include <mlir/Transforms/DialectConversion.h>
25#include <llvm/ADT/DenseMap.h>
26#include <llvm/ADT/StringSet.h>
27#include <llvm/ADT/Twine.h>
36inline mlir::DictionaryAttr
38 mlir::NamedAttrList newAttrs(attrs);
39 newAttrs.set(attrName, mlir::StringAttr::get(attrs.getContext(), name));
40 return newAttrs.getDictionary(attrs.getContext());
45inline mlir::DictionaryAttr
51inline mlir::DictionaryAttr
59 if (!usedNames.contains(desiredName)) {
60 usedNames.insert(desiredName);
61 return desiredName.str();
64 for (
unsigned suffix = 1;; ++suffix) {
65 std::string candidate = (desiredName +
"#" + llvm::Twine(suffix)).str();
66 if (!usedNames.contains(candidate)) {
67 usedNames.insert(candidate);
74inline std::optional<mlir::StringAttr>
76 if (!attrs || index >= attrs.size()) {
79 if (
auto dictAttr = llvm::dyn_cast<mlir::DictionaryAttr>(attrs[index])) {
80 if (
auto nameAttr = llvm::dyn_cast_if_present<mlir::StringAttr>(dictAttr.get(attrName))) {
95template <
typename GetNameAttrFn,
typename GetSplitSuffixesFn>
97 mlir::ArrayRef<mlir::Type> origTypes, GetNameAttrFn &&getNameAttr,
98 GetSplitSuffixesFn &&getSplitSuffixes
103 for (
auto [i, type] : llvm::enumerate(origTypes)) {
104 if (std::optional<mlir::StringAttr> nameAttr = getNameAttr(i)) {
118 mlir::ArrayAttr origAttrs,
const llvm::SmallVector<size_t> &originalIdxToSize,
119 const llvm::SmallVector<mlir::Type> &newTypes, llvm::StringRef functionNameAttrName,
120 llvm::ArrayRef<std::optional<llvm::StringRef>> origNames = {},
121 llvm::ArrayRef<llvm::StringRef> existingNames = {},
122 llvm::ArrayRef<llvm::SmallVector<std::string>> splitNameSuffixes = {}
127 assert(originalIdxToSize.size() == origAttrs.size());
128 if (originalIdxToSize.size() == newTypes.size()) {
132 llvm::SmallVector<mlir::Attribute> newAttrs;
133 llvm::StringSet<> usedNames;
134 if (!origNames.empty()) {
135 for (llvm::StringRef name : existingNames) {
136 usedNames.insert(name);
140 for (
auto [i, s] : llvm::enumerate(originalIdxToSize)) {
141 mlir::Attribute attr = origAttrs[i];
142 if (!origNames.empty() && !splitNameSuffixes.empty() && s != 1 && origNames[i]) {
143 assert(i < splitNameSuffixes.size());
144 assert(splitNameSuffixes[i].size() == s);
145 auto dictAttr = llvm::cast<mlir::DictionaryAttr>(attr);
146 llvm::StringRef name = *origNames[i];
147 for (llvm::StringRef suffix : splitNameSuffixes[i]) {
148 std::string desiredName = (llvm::Twine(name) + suffix).str();
155 newAttrs.append(s, attr);
157 return mlir::ArrayAttr::get(origAttrs.getContext(), newAttrs);
167 mlir::Location loc, mlir::TypeRange newResultTypes,
function::CallOp oldCall,
168 llvm::ArrayRef<mlir::ValueRange> mapOperands, mlir::ValueRange argOperands,
169 mlir::ConversionPatternRewriter &rewriter
171 llvm::SmallVector<mlir::Attribute> templateParams;
173 templateParams.append(templateParamsAttr.begin(), templateParamsAttr.end());
179 loc, newResultTypes, oldCall.
getCalleeAttr(), argOperands, templateParams
184 argOperands, templateParams
188 newCall->setDiscardableAttrs(oldCall->getDiscardableAttrDictionary());
193inline static mlir::Type replaceAffineMapArrayDimsWithWildcards(mlir::Type type) {
194 auto arrTy = llvm::dyn_cast<array::ArrayType>(type);
199 mlir::Builder builder(arrTy.getContext());
200 llvm::SmallVector<mlir::Attribute> dims;
201 dims.reserve(arrTy.getDimensionSizes().size());
202 for (mlir::Attribute dimSize : arrTy.getDimensionSizes()) {
203 if (llvm::isa<mlir::AffineMapAttr>(dimSize)) {
204 dims.push_back(builder.getIndexAttr(mlir::ShapedType::kDynamic));
206 dims.push_back(dimSize);
210 return arrTy.cloneWith(replaceAffineMapArrayDimsWithWildcards(arrTy.getElementType()), dims);
218 virtual llvm::SmallVector<mlir::Type>
convertInputs(mlir::ArrayRef<mlir::Type> origTypes) = 0;
219 virtual llvm::SmallVector<mlir::Type>
convertResults(mlir::ArrayRef<mlir::Type> origTypes) = 0;
221 virtual mlir::ArrayAttr
223 virtual mlir::ArrayAttr
234 llvm::SmallVector<mlir::Type> newInputs =
convertInputs(oldTy.getInputs());
235 llvm::SmallVector<mlir::Type> newResults =
convertResults(oldTy.getResults());
236 mlir::FunctionType newTy = mlir::FunctionType::get(
237 oldTy.getContext(), mlir::TypeRange(newInputs), mlir::TypeRange(newResults)
239 if (newTy == oldTy) {
246 rewriter.modifyOpInPlace(op, [&]() {
263 mlir::Block &entryBlock = body->front();
264 bool blockArgsNeedUpdate =
265 !std::cmp_equal(entryBlock.getNumArguments(), newInputs.size()) ||
266 llvm::any_of(llvm::zip_equal(entryBlock.getArgumentTypes(), newInputs), [](
auto pair) {
267 return std::get<0>(pair) != std::get<1>(pair);
269 if (blockArgsNeedUpdate) {
272 assert(std::cmp_equal(entryBlock.getNumArguments(), newInputs.size()));
273 for (
unsigned i = 0, e = entryBlock.getNumArguments(); i < e; ++i) {
274 assert(entryBlock.getArgument(i).getType() == newInputs[i]);
290 typename GenHeaderType,
typename IdType>
302 mlir::SymbolTableCollection &tables;
306 inline static void ensureImplementedAtCompile() {
308 sizeof(MemberRefOpClass) == 0,
309 "SplitAggregateInMemberRefOp not implemented for requested type."
318 static GenHeaderType
genHeader(MemberRefOpClass, mlir::ConversionPatternRewriter &) {
319 ensureImplementedAtCompile();
320 llvm_unreachable(
"must have concrete instantiation");
327 mlir::ConversionPatternRewriter &
329 ensureImplementedAtCompile();
330 llvm_unreachable(
"must have concrete instantiation");
337 mlir::MLIRContext *ctx, mlir::SymbolTableCollection &symTables,
340 :
mlir::OpConversionPattern<MemberRefOpClass>(ctx), tables(symTables),
341 repMapRef(memberRepMap) {}
343 static bool legal(MemberRefOpClass) {
344 ensureImplementedAtCompile();
345 llvm_unreachable(
"must have concrete instantiation");
350 MemberRefOpClass op,
OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter
352 if (ImplClass::legal(op)) {
353 return mlir::failure();
356 llvm::cast<component::MemberRefOpInterface>(op.getOperation()).getStructType();
359 assert(mlir::succeeded(tgtStructDef));
361 GenHeaderType prefixResult = ImplClass::genHeader(op, rewriter);
364 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
366 for (
const auto &[
id, newMember] : idToName) {
367 ImplClass::forId(op.getLoc(), prefixResult,
id, newMember, adaptor, rewriter);
369 if constexpr (
requires { ImplClass::finalize(op, prefixResult, adaptor, rewriter); }) {
370 ImplClass::finalize(op, prefixResult, adaptor, rewriter);
372 rewriter.eraseOp(op);
373 return mlir::success();
General helper for converting a FuncDefOp by changing its input and/or result types and the associate...
virtual void processBlockArgs(mlir::Block &entryBlock, mlir::RewriterBase &rewriter)=0
virtual llvm::SmallVector< mlir::Type > convertResults(mlir::ArrayRef< mlir::Type > origTypes)=0
virtual mlir::ArrayAttr convertResultAttrs(mlir::ArrayAttr origAttrs, llvm::SmallVector< mlir::Type > newTypes)=0
void convert(function::FuncDefOp op, mlir::RewriterBase &rewriter)
virtual ~FunctionTypeConverter()=default
virtual mlir::ArrayAttr convertInputAttrs(mlir::ArrayAttr origAttrs, llvm::SmallVector< mlir::Type > newTypes)=0
virtual llvm::SmallVector< mlir::Type > convertInputs(mlir::ArrayRef< mlir::Type > origTypes)=0
llvm::DenseMap< component::StructDefOp, llvm::DenseMap< mlir::StringAttr, LocalMemberReplacementMap > > MemberReplacementMap
Maps struct -> original aggregate-type member name -> LocalMemberReplacementMap.
static bool legal(MemberRefOpClass)
std::pair< mlir::StringAttr, mlir::Type > MemberInfo
Scalar member name and type.
SplitAggregateInMemberRefOp(mlir::MLIRContext *ctx, mlir::SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap)
typename MemberRefOpClass::Adaptor OpAdaptor
mlir::LogicalResult matchAndRewrite(MemberRefOpClass op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) const override
static GenHeaderType genHeader(MemberRefOpClass, mlir::ConversionPatternRewriter &)
Executed at the start of rewrite() to (optionally) generate anything that should appear before the pe...
static void forId(mlir::Location, GenHeaderType &, IdType, MemberInfo, OpAdaptor, mlir::ConversionPatternRewriter &)
Executed for each scalar id in the aggregate type of the original member to generate the per-scalar o...
llvm::DenseMap< IdType, MemberInfo > LocalMemberReplacementMap
Maps a scalar element identifier within the aggregate to its new scalar member info.
::mlir::FailureOr< SymbolLookupResult< StructDefOp > > getDefinition(::mlir::SymbolTableCollection &symbolTable, ::mlir::Operation *op, bool reportMissing=true) const
Gets the struct op that defines this struct.
::mlir::SymbolRefAttr getCalleeAttr()
::mlir::ArrayAttr getTemplateParamsAttr()
::mlir::OperandRangeRange getMapOperands()
::mlir::DenseI32ArrayAttr getNumDimsPerMapAttr()
::mlir::FunctionType getFunctionType()
void setArgAttrsAttr(::mlir::ArrayAttr attr)
void setResAttrsAttr(::mlir::ArrayAttr attr)
::mlir::ArrayAttr getArgAttrsAttr()
void setFunctionType(::mlir::FunctionType attrValue)
::mlir::Region * getCallableRegion()
Required by FunctionOpInterface.
::mlir::ArrayAttr getResAttrsAttr()
Restricts a template parameter to Op classes that implement the given OpInterface.
constexpr char ARG_NAME_ATTR_NAME[]
Attribute name for source-level function argument names.
constexpr char RES_NAME_ATTR_NAME[]
Attribute name for source-level function result names.
mlir::DictionaryAttr withFunctionResNameAttr(mlir::DictionaryAttr attrs, llvm::StringRef name)
Return a copy of the given result attribute dictionary with function.res_name set to name.
mlir::DictionaryAttr withFunctionNameAttr(mlir::DictionaryAttr attrs, llvm::StringRef attrName, llvm::StringRef name)
Return a copy of the given function argument/result attribute dictionary with attrName set to name.
mlir::ArrayAttr replicateFunctionNameAttrsAsNeeded(mlir::ArrayAttr origAttrs, const llvm::SmallVector< size_t > &originalIdxToSize, const llvm::SmallVector< mlir::Type > &newTypes, llvm::StringRef functionNameAttrName, llvm::ArrayRef< std::optional< llvm::StringRef > > origNames={}, llvm::ArrayRef< llvm::StringRef > existingNames={}, llvm::ArrayRef< llvm::SmallVector< std::string > > splitNameSuffixes={})
Expand function arg/result attribute arrays to match a split signature, rewriting name attrs with the...
mlir::DictionaryAttr withFunctionArgNameAttr(mlir::DictionaryAttr attrs, llvm::StringRef name)
Return a copy of the given argument attribute dictionary with function.arg_name set to name.
function::CallOp createCallPreservingInstantiationOperands(mlir::Location loc, mlir::TypeRange newResultTypes, function::CallOp oldCall, llvm::ArrayRef< mlir::ValueRange > mapOperands, mlir::ValueRange argOperands, mlir::ConversionPatternRewriter &rewriter)
Rebuild a function.call while preserving explicit instantiation state from oldCall.
SplitFunctionNameInfo collectSplitFunctionNameInfo(mlir::ArrayRef< mlir::Type > origTypes, GetNameAttrFn &&getNameAttr, GetSplitSuffixesFn &&getSplitSuffixes)
Collect function arg/result names and split suffixes from a list of original types.
std::optional< mlir::StringAttr > getAttrAtIndexWithName(mlir::ArrayAttr attrs, unsigned index, llvm::StringRef attrName)
Return the function arg/result attribute at index for the given name, if present.
std::string reserveUniqueAttrName(llvm::StringSet<> &usedNames, llvm::StringRef desiredName)
Reserve and return a unique function argument/result name based on desiredName.
Cached function arg/result names and split suffixes used while rewriting a function signature.
llvm::SmallVector< std::optional< llvm::StringRef > > originalNames
llvm::SmallVector< llvm::StringRef > existingNames
llvm::SmallVector< llvm::SmallVector< std::string > > splitNameSuffixes