10#include "llvm/IR/Constants.h"
18struct LoopMetadataConversion {
19 LoopMetadataConversion(
const llvm::MDNode *node, Location loc,
20 LoopAnnotationImporter &loopAnnotationImporter)
21 : node(node), loc(loc), loopAnnotationImporter(loopAnnotationImporter),
24 LoopAnnotationAttr convert();
27 LogicalResult initConversionState();
30 const llvm::MDNode *lookupAndEraseProperty(StringRef name);
35 FailureOr<BoolAttr> lookupUnitNode(StringRef name);
36 FailureOr<BoolAttr> lookupBoolNode(StringRef name,
bool negated =
false);
37 FailureOr<BoolAttr> lookupIntNodeAsBoolAttr(StringRef name);
38 FailureOr<IntegerAttr> lookupIntNode(StringRef name);
39 FailureOr<llvm::MDNode *> lookupMDNode(StringRef name);
40 FailureOr<SmallVector<llvm::MDNode *>> lookupMDNodes(StringRef name);
41 FailureOr<LoopAnnotationAttr> lookupFollowupNode(StringRef name);
42 FailureOr<BoolAttr> lookupBooleanUnitNode(StringRef enableName,
43 StringRef disableName,
44 bool negated =
false);
47 FailureOr<LoopVectorizeAttr> convertVectorizeAttr();
48 FailureOr<LoopInterleaveAttr> convertInterleaveAttr();
49 FailureOr<LoopUnrollAttr> convertUnrollAttr();
50 FailureOr<LoopUnrollAndJamAttr> convertUnrollAndJamAttr();
51 FailureOr<LoopLICMAttr> convertLICMAttr();
52 FailureOr<LoopDistributeAttr> convertDistributeAttr();
53 FailureOr<LoopPipelineAttr> convertPipelineAttr();
54 FailureOr<LoopPeeledAttr> convertPeeledAttr();
55 FailureOr<LoopUnswitchAttr> convertUnswitchAttr();
56 FailureOr<SmallVector<AccessGroupAttr>> convertParallelAccesses();
57 FusedLoc convertStartLoc();
58 FailureOr<FusedLoc> convertEndLoc();
60 llvm::SmallVector<llvm::DILocation *, 2> locations;
61 llvm::StringMap<const llvm::MDNode *> propertyMap;
62 const llvm::MDNode *node;
64 LoopAnnotationImporter &loopAnnotationImporter;
69LogicalResult LoopMetadataConversion::initConversionState() {
71 if (node->getNumOperands() == 0 ||
72 dyn_cast<llvm::MDNode>(node->getOperand(0)) != node)
75 for (
const llvm::MDOperand &operand : llvm::drop_begin(node->operands())) {
76 if (
auto *diLoc = dyn_cast<llvm::DILocation>(operand)) {
77 locations.push_back(diLoc);
81 auto *
property = dyn_cast<llvm::MDNode>(operand);
83 return emitWarning(loc) <<
"expected all loop properties to be either "
84 "debug locations or metadata nodes";
86 if (property->getNumOperands() == 0)
87 return emitWarning(loc) <<
"cannot import empty loop property";
89 auto *nameNode = dyn_cast<llvm::MDString>(property->getOperand(0));
91 return emitWarning(loc) <<
"cannot import loop property without a name";
92 StringRef name = nameNode->getString();
94 bool succ = propertyMap.try_emplace(name, property).second;
97 <<
"cannot import loop properties with duplicated names " << name;
104LoopMetadataConversion::lookupAndEraseProperty(StringRef name) {
105 auto it = propertyMap.find(name);
106 if (it == propertyMap.end())
108 const llvm::MDNode *
property = it->getValue();
109 propertyMap.erase(it);
113FailureOr<BoolAttr> LoopMetadataConversion::lookupUnitNode(StringRef name) {
114 const llvm::MDNode *
property = lookupAndEraseProperty(name);
116 return BoolAttr(
nullptr);
118 if (property->getNumOperands() != 1)
120 <<
"expected metadata node " << name <<
" to hold no value";
125FailureOr<BoolAttr> LoopMetadataConversion::lookupBooleanUnitNode(
126 StringRef enableName, StringRef disableName,
bool negated) {
127 auto enable = lookupUnitNode(enableName);
128 auto disable = lookupUnitNode(disableName);
132 if (*enable && *disable)
134 <<
"expected metadata nodes " << enableName <<
" and " << disableName
135 <<
" to be mutually exclusive.";
142 return BoolAttr(
nullptr);
145FailureOr<BoolAttr> LoopMetadataConversion::lookupBoolNode(StringRef name,
147 const llvm::MDNode *
property = lookupAndEraseProperty(name);
149 return BoolAttr(
nullptr);
151 auto emitNodeWarning = [&]() {
153 <<
"expected metadata node " << name <<
" to hold a boolean value";
156 if (property->getNumOperands() != 2)
157 return emitNodeWarning();
158 llvm::ConstantInt *val =
159 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
160 if (!val || val->getBitWidth() != 1)
161 return emitNodeWarning();
163 return BoolAttr::get(ctx, val->getValue().getLimitedValue(1) ^ negated);
167LoopMetadataConversion::lookupIntNodeAsBoolAttr(StringRef name) {
168 const llvm::MDNode *
property = lookupAndEraseProperty(name);
170 return BoolAttr(
nullptr);
172 auto emitNodeWarning = [&]() {
174 <<
"expected metadata node " << name <<
" to hold an integer value";
177 if (property->getNumOperands() != 2)
178 return emitNodeWarning();
179 llvm::ConstantInt *val =
180 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
181 if (!val || val->getBitWidth() != 32)
182 return emitNodeWarning();
184 return BoolAttr::get(ctx, val->getValue().getLimitedValue(1));
187FailureOr<IntegerAttr> LoopMetadataConversion::lookupIntNode(StringRef name) {
188 const llvm::MDNode *
property = lookupAndEraseProperty(name);
190 return IntegerAttr(
nullptr);
192 auto emitNodeWarning = [&]() {
194 <<
"expected metadata node " << name <<
" to hold an i32 value";
197 if (property->getNumOperands() != 2)
198 return emitNodeWarning();
200 llvm::ConstantInt *val =
201 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
202 if (!val || val->getBitWidth() != 32)
203 return emitNodeWarning();
205 return IntegerAttr::get(IntegerType::get(ctx, 32),
206 val->getValue().getLimitedValue());
209FailureOr<llvm::MDNode *> LoopMetadataConversion::lookupMDNode(StringRef name) {
210 const llvm::MDNode *
property = lookupAndEraseProperty(name);
214 auto emitNodeWarning = [&]() {
216 <<
"expected metadata node " << name <<
" to hold an MDNode";
219 if (property->getNumOperands() != 2)
220 return emitNodeWarning();
222 auto *node = dyn_cast<llvm::MDNode>(property->getOperand(1));
224 return emitNodeWarning();
229FailureOr<SmallVector<llvm::MDNode *>>
230LoopMetadataConversion::lookupMDNodes(StringRef name) {
231 const llvm::MDNode *
property = lookupAndEraseProperty(name);
232 SmallVector<llvm::MDNode *> res;
236 auto emitNodeWarning = [&]() {
237 return emitWarning(loc) <<
"expected metadata node " << name
238 <<
" to hold one or multiple MDNodes";
241 if (property->getNumOperands() < 2)
242 return emitNodeWarning();
244 for (
unsigned i = 1, e = property->getNumOperands(); i < e; ++i) {
245 auto *node = dyn_cast<llvm::MDNode>(property->getOperand(i));
247 return emitNodeWarning();
254FailureOr<LoopAnnotationAttr>
255LoopMetadataConversion::lookupFollowupNode(StringRef name) {
256 auto node = lookupMDNode(name);
259 if (*node ==
nullptr)
260 return LoopAnnotationAttr(
nullptr);
274template <
typename T,
typename... P>
276 bool anyFailed = (failed(args) || ...);
284 return T::get(ctx, *args...);
287FailureOr<LoopVectorizeAttr> LoopMetadataConversion::convertVectorizeAttr() {
288 FailureOr<BoolAttr> enable = lookupBooleanUnitNode(
289 "llvm.loop.vectorize.enable",
"llvm.loop.vectorize.disable",
291 FailureOr<BoolAttr> predicateEnable =
292 lookupBooleanUnitNode(
"llvm.loop.vectorize.predicate.enable",
293 "llvm.loop.vectorize.predicate.disable");
294 FailureOr<BoolAttr> scalableEnable =
295 lookupBooleanUnitNode(
"llvm.loop.vectorize.scalable.enable",
296 "llvm.loop.vectorize.scalable.disable");
297 FailureOr<IntegerAttr> width = lookupIntNode(
"llvm.loop.vectorize.width");
298 FailureOr<LoopAnnotationAttr> followupVec =
299 lookupFollowupNode(
"llvm.loop.vectorize.followup_vectorized");
300 FailureOr<LoopAnnotationAttr> followupEpi =
301 lookupFollowupNode(
"llvm.loop.vectorize.followup_epilogue");
302 FailureOr<LoopAnnotationAttr> followupAll =
303 lookupFollowupNode(
"llvm.loop.vectorize.followup_all");
306 scalableEnable, width, followupVec,
307 followupEpi, followupAll);
310FailureOr<LoopInterleaveAttr> LoopMetadataConversion::convertInterleaveAttr() {
311 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.interleave.count");
315FailureOr<LoopUnrollAttr> LoopMetadataConversion::convertUnrollAttr() {
316 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
317 "llvm.loop.unroll.enable",
"llvm.loop.unroll.disable",
true);
318 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.unroll.count");
319 FailureOr<BoolAttr> runtimeDisable =
320 lookupUnitNode(
"llvm.loop.unroll.runtime.disable");
321 FailureOr<BoolAttr> full = lookupUnitNode(
"llvm.loop.unroll.full");
322 FailureOr<LoopAnnotationAttr> followupUnrolled =
323 lookupFollowupNode(
"llvm.loop.unroll.followup_unrolled");
324 FailureOr<LoopAnnotationAttr> followupRemainder =
325 lookupFollowupNode(
"llvm.loop.unroll.followup_remainder");
326 FailureOr<LoopAnnotationAttr> followupAll =
327 lookupFollowupNode(
"llvm.loop.unroll.followup_all");
330 full, followupUnrolled,
331 followupRemainder, followupAll);
334FailureOr<LoopUnrollAndJamAttr>
335LoopMetadataConversion::convertUnrollAndJamAttr() {
336 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
337 "llvm.loop.unroll_and_jam.enable",
"llvm.loop.unroll_and_jam.disable",
339 FailureOr<IntegerAttr> count =
340 lookupIntNode(
"llvm.loop.unroll_and_jam.count");
341 FailureOr<LoopAnnotationAttr> followupOuter =
342 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_outer");
343 FailureOr<LoopAnnotationAttr> followupInner =
344 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_inner");
345 FailureOr<LoopAnnotationAttr> followupRemainderOuter =
346 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_remainder_outer");
347 FailureOr<LoopAnnotationAttr> followupRemainderInner =
348 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_remainder_inner");
349 FailureOr<LoopAnnotationAttr> followupAll =
350 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_all");
352 ctx, disable, count, followupOuter, followupInner, followupRemainderOuter,
353 followupRemainderInner, followupAll);
356FailureOr<LoopLICMAttr> LoopMetadataConversion::convertLICMAttr() {
357 FailureOr<BoolAttr> disable = lookupUnitNode(
"llvm.licm.disable");
358 FailureOr<BoolAttr> versioningDisable =
359 lookupUnitNode(
"llvm.loop.licm_versioning.disable");
363FailureOr<LoopDistributeAttr> LoopMetadataConversion::convertDistributeAttr() {
364 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
365 "llvm.loop.distribute.enable",
"llvm.loop.distribute.disable",
367 FailureOr<LoopAnnotationAttr> followupCoincident =
368 lookupFollowupNode(
"llvm.loop.distribute.followup_coincident");
369 FailureOr<LoopAnnotationAttr> followupSequential =
370 lookupFollowupNode(
"llvm.loop.distribute.followup_sequential");
371 FailureOr<LoopAnnotationAttr> followupFallback =
372 lookupFollowupNode(
"llvm.loop.distribute.followup_fallback");
373 FailureOr<LoopAnnotationAttr> followupAll =
374 lookupFollowupNode(
"llvm.loop.distribute.followup_all");
377 followupFallback, followupAll);
380FailureOr<LoopPipelineAttr> LoopMetadataConversion::convertPipelineAttr() {
381 FailureOr<BoolAttr> disable = lookupBoolNode(
"llvm.loop.pipeline.disable");
382 FailureOr<IntegerAttr> initiationinterval =
383 lookupIntNode(
"llvm.loop.pipeline.initiationinterval");
387FailureOr<LoopPeeledAttr> LoopMetadataConversion::convertPeeledAttr() {
388 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.peeled.count");
392FailureOr<LoopUnswitchAttr> LoopMetadataConversion::convertUnswitchAttr() {
393 FailureOr<BoolAttr> partialDisable =
394 lookupUnitNode(
"llvm.loop.unswitch.partial.disable");
398FailureOr<SmallVector<AccessGroupAttr>>
399LoopMetadataConversion::convertParallelAccesses() {
400 FailureOr<SmallVector<llvm::MDNode *>> nodes =
401 lookupMDNodes(
"llvm.loop.parallel_accesses");
404 SmallVector<AccessGroupAttr> refs;
405 for (llvm::MDNode *node : *nodes) {
406 FailureOr<SmallVector<AccessGroupAttr>> accessGroups =
408 if (
failed(accessGroups)) {
409 emitWarning(loc) <<
"could not lookup access group";
412 llvm::append_range(refs, *accessGroups);
417FusedLoc LoopMetadataConversion::convertStartLoc() {
418 if (locations.empty())
420 return dyn_cast<FusedLoc>(
424FailureOr<FusedLoc> LoopMetadataConversion::convertEndLoc() {
425 if (locations.size() < 2)
427 if (locations.size() > 2)
429 <<
"expected loop metadata to have at most two DILocations";
430 return dyn_cast<FusedLoc>(
434LoopAnnotationAttr LoopMetadataConversion::convert() {
435 if (
failed(initConversionState()))
438 FailureOr<BoolAttr> disableNonForced =
439 lookupUnitNode(
"llvm.loop.disable_nonforced");
440 FailureOr<LoopVectorizeAttr> vecAttr = convertVectorizeAttr();
441 FailureOr<LoopInterleaveAttr> interleaveAttr = convertInterleaveAttr();
442 FailureOr<LoopUnrollAttr> unrollAttr = convertUnrollAttr();
443 FailureOr<LoopUnrollAndJamAttr> unrollAndJamAttr = convertUnrollAndJamAttr();
444 FailureOr<LoopLICMAttr> licmAttr = convertLICMAttr();
445 FailureOr<LoopDistributeAttr> distributeAttr = convertDistributeAttr();
446 FailureOr<LoopPipelineAttr> pipelineAttr = convertPipelineAttr();
447 FailureOr<LoopPeeledAttr> peeledAttr = convertPeeledAttr();
448 FailureOr<LoopUnswitchAttr> unswitchAttr = convertUnswitchAttr();
449 FailureOr<BoolAttr> mustProgress = lookupUnitNode(
"llvm.loop.mustprogress");
450 FailureOr<BoolAttr> isVectorized =
451 lookupIntNodeAsBoolAttr(
"llvm.loop.isvectorized");
452 FailureOr<SmallVector<AccessGroupAttr>> parallelAccesses =
453 convertParallelAccesses();
456 if (!propertyMap.empty()) {
457 for (
auto name : propertyMap.keys())
458 emitWarning(loc) <<
"unknown loop annotation " << name;
462 FailureOr<FusedLoc> startLoc = convertStartLoc();
463 FailureOr<FusedLoc> endLoc = convertEndLoc();
466 ctx, disableNonForced, vecAttr, interleaveAttr, unrollAttr,
467 unrollAndJamAttr, licmAttr, distributeAttr, pipelineAttr, peeledAttr,
468 unswitchAttr, mustProgress, isVectorized, startLoc, endLoc,
480 auto it = loopMetadataMapping.find(node);
481 if (it != loopMetadataMapping.end())
482 return it->getSecond();
484 LoopAnnotationAttr attr = LoopMetadataConversion(node, loc, *
this).convert();
486 mapLoopMetadata(node, attr);
494 if (!node->getNumOperands())
495 accessGroups.push_back(node);
496 for (
const llvm::MDOperand &operand : node->operands()) {
497 auto *childNode = dyn_cast<llvm::MDNode>(operand);
500 accessGroups.push_back(cast<llvm::MDNode>(operand.get()));
504 for (
const llvm::MDNode *accessGroup : accessGroups) {
505 if (accessGroupMapping.count(accessGroup))
508 if (accessGroup->getNumOperands() != 0 || !accessGroup->isDistinct())
510 <<
"expected an access group node to be empty and distinct";
513 accessGroupMapping[accessGroup] = builder.getAttr<AccessGroupAttr>();
518FailureOr<SmallVector<AccessGroupAttr>>
523 if (!node->getNumOperands())
524 accessGroups.push_back(accessGroupMapping.lookup(node));
525 for (
const llvm::MDOperand &operand : node->operands()) {
526 auto *node = cast<llvm::MDNode>(operand.get());
527 accessGroups.push_back(accessGroupMapping.lookup(node));
530 if (llvm::is_contained(accessGroups,
nullptr))
static T createIfNonNull(MLIRContext *ctx, const P &...args)
Helper function that only creates and attribute of type T if all argument conversion were successfull...
static bool isEmptyOrNull(const Attribute attr)
Attributes are known-constant values of operations.
static BoolAttr get(MLIRContext *context, bool value)
Location translateLoc(llvm::DILocation *loc)
Translates the debug location.
LoopAnnotationAttr translateLoopAnnotation(const llvm::MDNode *node, Location loc)
LogicalResult translateAccessGroup(const llvm::MDNode *node, Location loc)
Converts all LLVM access groups starting from node to MLIR access group attributes.
ModuleImport & moduleImport
The ModuleImport owning this instance.
FailureOr< SmallVector< AccessGroupAttr > > lookupAccessGroupAttrs(const llvm::MDNode *node) const
Returns the access group attribute that map to the access group nodes starting from the access group ...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
Include the generated interface declarations.
InFlightDiagnostic emitWarning(Location loc)
Utility method to emit a warning message using this location.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.