10#include "llvm/ADT/ScopeExit.h"
11#include "llvm/ADT/SmallPtrSet.h"
12#include "llvm/IR/Constants.h"
20struct LoopMetadataConversion {
22 static LoopAnnotationAttr convert(
const llvm::MDNode *node, Location loc,
23 LoopAnnotationImporter &importer);
26 LoopMetadataConversion(
27 const llvm::MDNode *node, Location loc,
28 LoopAnnotationImporter &loopAnnotationImporter,
29 llvm::SmallPtrSetImpl<const llvm::MDNode *> &activeNodes)
30 : node(node), loc(loc), loopAnnotationImporter(loopAnnotationImporter),
31 ctx(loc->
getContext()), activeNodes(activeNodes) {}
34 LoopAnnotationAttr convertFollowup(
const llvm::MDNode *followup);
37 LoopAnnotationAttr convertProperties();
40 LogicalResult initConversionState();
43 const llvm::MDNode *lookupAndEraseProperty(StringRef name);
48 FailureOr<BoolAttr> lookupUnitNode(StringRef name);
49 FailureOr<BoolAttr> lookupBoolNode(StringRef name,
bool negated =
false);
50 FailureOr<BoolAttr> lookupIntNodeAsBoolAttr(StringRef name);
51 FailureOr<IntegerAttr> lookupIntNode(StringRef name);
52 FailureOr<SmallVector<llvm::MDNode *>> lookupMDNodes(StringRef name);
53 FailureOr<LoopAnnotationAttr> lookupFollowupNode(StringRef name);
54 FailureOr<BoolAttr> lookupBooleanUnitNode(StringRef enableName,
55 StringRef disableName,
56 bool negated =
false);
59 FailureOr<LoopVectorizeAttr> convertVectorizeAttr();
60 FailureOr<LoopInterleaveAttr> convertInterleaveAttr();
61 FailureOr<LoopUnrollAttr> convertUnrollAttr();
62 FailureOr<LoopUnrollAndJamAttr> convertUnrollAndJamAttr();
63 FailureOr<LoopLICMAttr> convertLICMAttr();
64 FailureOr<LoopDistributeAttr> convertDistributeAttr();
65 FailureOr<LoopPipelineAttr> convertPipelineAttr();
66 FailureOr<LoopPeeledAttr> convertPeeledAttr();
67 FailureOr<LoopUnswitchAttr> convertUnswitchAttr();
68 FailureOr<SmallVector<AccessGroupAttr>> convertParallelAccesses();
69 FusedLoc convertStartLoc();
70 FailureOr<FusedLoc> convertEndLoc();
72 llvm::SmallVector<llvm::DILocation *, 2> locations;
73 llvm::StringMap<const llvm::MDNode *> propertyMap;
74 const llvm::MDNode *node;
76 LoopAnnotationImporter &loopAnnotationImporter;
78 llvm::SmallPtrSetImpl<const llvm::MDNode *> &activeNodes;
82LogicalResult LoopMetadataConversion::initConversionState() {
83 for (
const llvm::MDOperand &operand : llvm::drop_begin(node->operands())) {
84 if (
auto *diLoc = dyn_cast<llvm::DILocation>(operand)) {
85 locations.push_back(diLoc);
89 auto *
property = dyn_cast<llvm::MDNode>(operand);
91 return emitWarning(loc) <<
"expected all loop properties to be either "
92 "debug locations or metadata nodes";
94 if (property->getNumOperands() == 0)
95 return emitWarning(loc) <<
"cannot import empty loop property";
97 auto *nameNode = dyn_cast<llvm::MDString>(property->getOperand(0));
99 return emitWarning(loc) <<
"cannot import loop property without a name";
100 StringRef name = nameNode->getString();
102 bool succ = propertyMap.try_emplace(name, property).second;
105 <<
"cannot import loop properties with duplicated names " << name;
112LoopMetadataConversion::lookupAndEraseProperty(StringRef name) {
113 auto it = propertyMap.find(name);
114 if (it == propertyMap.end())
116 const llvm::MDNode *
property = it->getValue();
117 propertyMap.erase(it);
121FailureOr<BoolAttr> LoopMetadataConversion::lookupUnitNode(StringRef name) {
122 const llvm::MDNode *
property = lookupAndEraseProperty(name);
124 return BoolAttr(
nullptr);
126 if (property->getNumOperands() != 1)
128 <<
"expected metadata node " << name <<
" to hold no value";
133FailureOr<BoolAttr> LoopMetadataConversion::lookupBooleanUnitNode(
134 StringRef enableName, StringRef disableName,
bool negated) {
135 auto enable = lookupUnitNode(enableName);
136 auto disable = lookupUnitNode(disableName);
140 if (*enable && *disable)
142 <<
"expected metadata nodes " << enableName <<
" and " << disableName
143 <<
" to be mutually exclusive.";
150 return BoolAttr(
nullptr);
153FailureOr<BoolAttr> LoopMetadataConversion::lookupBoolNode(StringRef name,
155 const llvm::MDNode *
property = lookupAndEraseProperty(name);
157 return BoolAttr(
nullptr);
159 auto emitNodeWarning = [&]() {
161 <<
"expected metadata node " << name <<
" to hold a boolean value";
164 if (property->getNumOperands() != 2)
165 return emitNodeWarning();
166 llvm::ConstantInt *val =
167 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
168 if (!val || val->getBitWidth() != 1)
169 return emitNodeWarning();
171 return BoolAttr::get(ctx, val->getValue().getLimitedValue(1) ^ negated);
175LoopMetadataConversion::lookupIntNodeAsBoolAttr(StringRef name) {
176 const llvm::MDNode *
property = lookupAndEraseProperty(name);
178 return BoolAttr(
nullptr);
180 auto emitNodeWarning = [&]() {
182 <<
"expected metadata node " << name <<
" to hold an integer value";
185 if (property->getNumOperands() != 2)
186 return emitNodeWarning();
187 llvm::ConstantInt *val =
188 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
189 if (!val || val->getBitWidth() != 32)
190 return emitNodeWarning();
192 return BoolAttr::get(ctx, val->getValue().getLimitedValue(1));
195FailureOr<IntegerAttr> LoopMetadataConversion::lookupIntNode(StringRef name) {
196 const llvm::MDNode *
property = lookupAndEraseProperty(name);
198 return IntegerAttr(
nullptr);
200 auto emitNodeWarning = [&]() {
202 <<
"expected metadata node " << name <<
" to hold an i32 value";
205 if (property->getNumOperands() != 2)
206 return emitNodeWarning();
208 llvm::ConstantInt *val =
209 llvm::mdconst::dyn_extract<llvm::ConstantInt>(property->getOperand(1));
210 if (!val || val->getBitWidth() != 32)
211 return emitNodeWarning();
213 return IntegerAttr::get(IntegerType::get(ctx, 32),
214 val->getValue().getLimitedValue());
217FailureOr<SmallVector<llvm::MDNode *>>
218LoopMetadataConversion::lookupMDNodes(StringRef name) {
219 const llvm::MDNode *
property = lookupAndEraseProperty(name);
220 SmallVector<llvm::MDNode *> res;
224 auto emitNodeWarning = [&]() {
225 return emitWarning(loc) <<
"expected metadata node " << name
226 <<
" to hold one or multiple MDNodes";
229 if (property->getNumOperands() < 2)
230 return emitNodeWarning();
232 for (
unsigned i = 1, e = property->getNumOperands(); i < e; ++i) {
233 auto *node = dyn_cast<llvm::MDNode>(property->getOperand(i));
235 return emitNodeWarning();
242FailureOr<LoopAnnotationAttr>
243LoopMetadataConversion::lookupFollowupNode(StringRef name) {
244 const llvm::MDNode *followup = lookupAndEraseProperty(name);
246 return LoopAnnotationAttr(
nullptr);
248 LoopAnnotationAttr attr = convertFollowup(followup);
263template <
typename T,
typename... P>
265 bool anyFailed = (failed(args) || ...);
273 return T::get(ctx, *args...);
276FailureOr<LoopVectorizeAttr> LoopMetadataConversion::convertVectorizeAttr() {
277 FailureOr<BoolAttr> enable = lookupBooleanUnitNode(
278 "llvm.loop.vectorize.enable",
"llvm.loop.vectorize.disable",
280 FailureOr<BoolAttr> predicateEnable =
281 lookupBooleanUnitNode(
"llvm.loop.vectorize.predicate.enable",
282 "llvm.loop.vectorize.predicate.disable");
283 FailureOr<BoolAttr> scalableEnable =
284 lookupBooleanUnitNode(
"llvm.loop.vectorize.scalable.enable",
285 "llvm.loop.vectorize.scalable.disable");
286 FailureOr<IntegerAttr> width = lookupIntNode(
"llvm.loop.vectorize.width");
287 FailureOr<LoopAnnotationAttr> followupVec =
288 lookupFollowupNode(
"llvm.loop.vectorize.followup_vectorized");
289 FailureOr<LoopAnnotationAttr> followupEpi =
290 lookupFollowupNode(
"llvm.loop.vectorize.followup_epilogue");
291 FailureOr<LoopAnnotationAttr> followupAll =
292 lookupFollowupNode(
"llvm.loop.vectorize.followup_all");
295 scalableEnable, width, followupVec,
296 followupEpi, followupAll);
299FailureOr<LoopInterleaveAttr> LoopMetadataConversion::convertInterleaveAttr() {
300 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.interleave.count");
304FailureOr<LoopUnrollAttr> LoopMetadataConversion::convertUnrollAttr() {
305 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
306 "llvm.loop.unroll.enable",
"llvm.loop.unroll.disable",
true);
307 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.unroll.count");
308 FailureOr<BoolAttr> runtimeDisable =
309 lookupUnitNode(
"llvm.loop.unroll.runtime.disable");
310 FailureOr<BoolAttr> full = lookupUnitNode(
"llvm.loop.unroll.full");
311 FailureOr<LoopAnnotationAttr> followupUnrolled =
312 lookupFollowupNode(
"llvm.loop.unroll.followup_unrolled");
313 FailureOr<LoopAnnotationAttr> followupRemainder =
314 lookupFollowupNode(
"llvm.loop.unroll.followup_remainder");
315 FailureOr<LoopAnnotationAttr> followupAll =
316 lookupFollowupNode(
"llvm.loop.unroll.followup_all");
319 full, followupUnrolled,
320 followupRemainder, followupAll);
323FailureOr<LoopUnrollAndJamAttr>
324LoopMetadataConversion::convertUnrollAndJamAttr() {
325 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
326 "llvm.loop.unroll_and_jam.enable",
"llvm.loop.unroll_and_jam.disable",
328 FailureOr<IntegerAttr> count =
329 lookupIntNode(
"llvm.loop.unroll_and_jam.count");
330 FailureOr<LoopAnnotationAttr> followupOuter =
331 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_outer");
332 FailureOr<LoopAnnotationAttr> followupInner =
333 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_inner");
334 FailureOr<LoopAnnotationAttr> followupRemainderOuter =
335 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_remainder_outer");
336 FailureOr<LoopAnnotationAttr> followupRemainderInner =
337 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_remainder_inner");
338 FailureOr<LoopAnnotationAttr> followupAll =
339 lookupFollowupNode(
"llvm.loop.unroll_and_jam.followup_all");
341 ctx, disable, count, followupOuter, followupInner, followupRemainderOuter,
342 followupRemainderInner, followupAll);
345FailureOr<LoopLICMAttr> LoopMetadataConversion::convertLICMAttr() {
346 FailureOr<BoolAttr> disable = lookupUnitNode(
"llvm.licm.disable");
347 FailureOr<BoolAttr> versioningDisable =
348 lookupUnitNode(
"llvm.loop.licm_versioning.disable");
352FailureOr<LoopDistributeAttr> LoopMetadataConversion::convertDistributeAttr() {
353 FailureOr<BoolAttr> disable = lookupBooleanUnitNode(
354 "llvm.loop.distribute.enable",
"llvm.loop.distribute.disable",
356 FailureOr<LoopAnnotationAttr> followupCoincident =
357 lookupFollowupNode(
"llvm.loop.distribute.followup_coincident");
358 FailureOr<LoopAnnotationAttr> followupSequential =
359 lookupFollowupNode(
"llvm.loop.distribute.followup_sequential");
360 FailureOr<LoopAnnotationAttr> followupFallback =
361 lookupFollowupNode(
"llvm.loop.distribute.followup_fallback");
362 FailureOr<LoopAnnotationAttr> followupAll =
363 lookupFollowupNode(
"llvm.loop.distribute.followup_all");
366 followupFallback, followupAll);
369FailureOr<LoopPipelineAttr> LoopMetadataConversion::convertPipelineAttr() {
370 FailureOr<BoolAttr> disable = lookupBoolNode(
"llvm.loop.pipeline.disable");
371 FailureOr<IntegerAttr> initiationinterval =
372 lookupIntNode(
"llvm.loop.pipeline.initiationinterval");
376FailureOr<LoopPeeledAttr> LoopMetadataConversion::convertPeeledAttr() {
377 FailureOr<IntegerAttr> count = lookupIntNode(
"llvm.loop.peeled.count");
381FailureOr<LoopUnswitchAttr> LoopMetadataConversion::convertUnswitchAttr() {
382 FailureOr<BoolAttr> partialDisable =
383 lookupUnitNode(
"llvm.loop.unswitch.partial.disable");
387FailureOr<SmallVector<AccessGroupAttr>>
388LoopMetadataConversion::convertParallelAccesses() {
389 FailureOr<SmallVector<llvm::MDNode *>> nodes =
390 lookupMDNodes(
"llvm.loop.parallel_accesses");
393 SmallVector<AccessGroupAttr> refs;
394 for (llvm::MDNode *node : *nodes) {
395 FailureOr<SmallVector<AccessGroupAttr>> accessGroups =
397 if (
failed(accessGroups)) {
398 emitWarning(loc) <<
"could not lookup access group";
401 llvm::append_range(refs, *accessGroups);
406FusedLoc LoopMetadataConversion::convertStartLoc() {
407 if (locations.empty())
409 return dyn_cast<FusedLoc>(
413FailureOr<FusedLoc> LoopMetadataConversion::convertEndLoc() {
414 if (locations.size() < 2)
416 if (locations.size() > 2)
418 <<
"expected loop metadata to have at most two DILocations";
419 return dyn_cast<FusedLoc>(
424LoopMetadataConversion::convert(
const llvm::MDNode *node, Location loc,
425 LoopAnnotationImporter &importer) {
426 if (node->getNumOperands() == 0 ||
427 dyn_cast<llvm::MDNode>(node->getOperand(0)) != node) {
432 llvm::SmallPtrSet<const llvm::MDNode *, 4> activeNodes;
433 return LoopMetadataConversion(node, loc, importer, activeNodes)
434 .convertProperties();
438LoopMetadataConversion::convertFollowup(
const llvm::MDNode *followup) {
439 if (followup->getNumOperands() == 0 ||
440 !isa<llvm::MDString>(followup->getOperand(0))) {
446 if (followup->getNumOperands() == 1)
447 return LoopAnnotationAttr::get(ctx, {}, {}, {}, {}, {}, {}, {}, {}, {}, {},
450 return LoopMetadataConversion(followup, loc, loopAnnotationImporter,
452 .convertProperties();
455LoopAnnotationAttr LoopMetadataConversion::convertProperties() {
456 if (!activeNodes.insert(node).second) {
457 emitWarning(loc) <<
"cannot import cyclic loop annotation";
460 llvm::scope_exit guard([&] { activeNodes.erase(node); });
462 if (
failed(initConversionState()))
465 FailureOr<BoolAttr> disableNonForced =
466 lookupUnitNode(
"llvm.loop.disable_nonforced");
467 FailureOr<LoopVectorizeAttr> vecAttr = convertVectorizeAttr();
468 FailureOr<LoopInterleaveAttr> interleaveAttr = convertInterleaveAttr();
469 FailureOr<LoopUnrollAttr> unrollAttr = convertUnrollAttr();
470 FailureOr<LoopUnrollAndJamAttr> unrollAndJamAttr = convertUnrollAndJamAttr();
471 FailureOr<LoopLICMAttr> licmAttr = convertLICMAttr();
472 FailureOr<LoopDistributeAttr> distributeAttr = convertDistributeAttr();
473 FailureOr<LoopPipelineAttr> pipelineAttr = convertPipelineAttr();
474 FailureOr<LoopPeeledAttr> peeledAttr = convertPeeledAttr();
475 FailureOr<LoopUnswitchAttr> unswitchAttr = convertUnswitchAttr();
476 FailureOr<BoolAttr> mustProgress = lookupUnitNode(
"llvm.loop.mustprogress");
477 FailureOr<BoolAttr> isVectorized =
478 lookupIntNodeAsBoolAttr(
"llvm.loop.isvectorized");
479 FailureOr<SmallVector<AccessGroupAttr>> parallelAccesses =
480 convertParallelAccesses();
483 if (!propertyMap.empty()) {
484 for (
auto name : propertyMap.keys())
485 emitWarning(loc) <<
"unknown loop annotation " << name;
489 FailureOr<FusedLoc> startLoc = convertStartLoc();
490 FailureOr<FusedLoc> endLoc = convertEndLoc();
493 ctx, disableNonForced, vecAttr, interleaveAttr, unrollAttr,
494 unrollAndJamAttr, licmAttr, distributeAttr, pipelineAttr, peeledAttr,
495 unswitchAttr, mustProgress, isVectorized, startLoc, endLoc,
507 auto it = loopMetadataMapping.find(node);
508 if (it != loopMetadataMapping.end())
509 return it->getSecond();
511 LoopAnnotationAttr attr = LoopMetadataConversion::convert(node, loc, *
this);
513 mapLoopMetadata(node, attr);
521 if (!node->getNumOperands())
522 accessGroups.push_back(node);
523 for (
const llvm::MDOperand &operand : node->operands()) {
524 auto *childNode = dyn_cast<llvm::MDNode>(operand);
527 accessGroups.push_back(cast<llvm::MDNode>(operand.get()));
531 for (
const llvm::MDNode *accessGroup : accessGroups) {
532 if (accessGroupMapping.count(accessGroup))
535 if (accessGroup->getNumOperands() != 0 || !accessGroup->isDistinct())
537 <<
"expected an access group node to be empty and distinct";
540 accessGroupMapping[accessGroup] = builder.getAttr<AccessGroupAttr>();
545FailureOr<SmallVector<AccessGroupAttr>>
550 if (!node->getNumOperands())
551 accessGroups.push_back(accessGroupMapping.lookup(node));
552 for (
const llvm::MDOperand &operand : node->operands()) {
553 auto *node = cast<llvm::MDNode>(operand.get());
554 accessGroups.push_back(accessGroupMapping.lookup(node));
557 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.