28#include "llvm/ADT/ArrayRef.h"
29#include "llvm/ADT/PostOrderIterator.h"
30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/STLForwardCompat.h"
32#include "llvm/ADT/SmallString.h"
33#include "llvm/ADT/StringExtras.h"
34#include "llvm/ADT/StringRef.h"
35#include "llvm/ADT/TypeSwitch.h"
36#include "llvm/ADT/bit.h"
37#include "llvm/Support/InterleavedRange.h"
43#include "mlir/Dialect/OpenMP/OpenMPOpsDialect.cpp.inc"
44#include "mlir/Dialect/OpenMP/OpenMPOpsEnums.cpp.inc"
45#include "mlir/Dialect/OpenMP/OpenMPOpsInterfaces.cpp.inc"
46#include "mlir/Dialect/OpenMP/OpenMPTypeInterfaces.cpp.inc"
53 return attrs.empty() ?
nullptr : ArrayAttr::get(context, attrs);
67struct MemRefPointerLikeModel
68 :
public PointerLikeType::ExternalModel<MemRefPointerLikeModel,
71 return llvm::cast<MemRefType>(pointer).getElementType();
75struct LLVMPointerPointerLikeModel
76 :
public PointerLikeType::ExternalModel<LLVMPointerPointerLikeModel,
77 LLVM::LLVMPointerType> {
102 bool isRegionArgOfOp;
112 assert(isRegionArgOfOp &&
"Must describe a region operand");
115 size_t &getArgIdx() {
116 assert(isRegionArgOfOp &&
"Must describe a region operand");
121 assert(!isRegionArgOfOp &&
"Must describe a operation of a region");
125 assert(!isRegionArgOfOp &&
"Must describe a operation of a region");
128 bool isLoopOp()
const {
129 assert(!isRegionArgOfOp &&
"Must describe a operation of a region");
130 return isa<CanonicalLoopOp>(op);
132 Region *&getParentRegion() {
133 assert(!isRegionArgOfOp &&
"Must describe a operation of a region");
136 size_t &getLoopDepth() {
137 assert(!isRegionArgOfOp &&
"Must describe a operation of a region");
141 void skipIf(
bool v =
true) { skip = skip || v; }
159 llvm::ReversePostOrderTraversal<Block *> traversal(&r->
getBlocks().front());
162 size_t sequentialIdx = -1;
163 bool isOnlyContainerOp =
true;
164 for (
Block *
b : traversal) {
166 if (&op == o && !found) {
170 if (op.getNumRegions()) {
173 isOnlyContainerOp =
false;
175 if (found && !isOnlyContainerOp)
180 Component &containerOpInRegion = components.emplace_back();
181 containerOpInRegion.isRegionArgOfOp =
false;
182 containerOpInRegion.isUnique = isOnlyContainerOp;
183 containerOpInRegion.getContainerOp() = o;
184 containerOpInRegion.getOpPos() = sequentialIdx;
185 containerOpInRegion.getParentRegion() = r;
190 Component ®ionArgOfOperation = components.emplace_back();
191 regionArgOfOperation.isRegionArgOfOp =
true;
192 regionArgOfOperation.isUnique =
true;
193 regionArgOfOperation.getArgIdx() = 0;
194 regionArgOfOperation.getOwnerOp() = parent;
206 for (
auto [idx, region] : llvm::enumerate(o->
getRegions())) {
210 llvm_unreachable(
"Region not child of its parent operation");
212 regionArgOfOperation.isUnique =
false;
213 regionArgOfOperation.getArgIdx() = getRegionIndex(parent, r);
221 for (Component &c : components)
222 c.skipIf(c.isRegionArgOfOp && c.isUnique);
225 size_t numSurroundingLoops = 0;
226 for (Component &c : llvm::reverse(components)) {
231 if (c.isRegionArgOfOp) {
232 numSurroundingLoops = 0;
239 numSurroundingLoops = 0;
241 c.getLoopDepth() = numSurroundingLoops;
244 if (isa<CanonicalLoopOp>(c.getContainerOp()))
245 numSurroundingLoops += 1;
250 bool isLoopNest =
false;
251 for (Component &c : components) {
252 if (c.skip || c.isRegionArgOfOp)
255 if (!isLoopNest && c.getLoopDepth() >= 1) {
258 }
else if (isLoopNest) {
260 c.skipIf(c.isUnique);
264 if (c.getLoopDepth() == 0)
271 for (Component &c : components)
272 c.skipIf(!c.isRegionArgOfOp && c.isUnique &&
273 !isa<CanonicalLoopOp>(c.getContainerOp()));
277 bool newRegion =
true;
278 for (Component &c : llvm::reverse(components)) {
279 c.skipIf(newRegion && c.isUnique);
286 if (!c.isRegionArgOfOp && c.getContainerOp())
292 llvm::raw_svector_ostream NameOS(Name);
293 for (
auto &c : llvm::reverse(components)) {
297 if (c.isRegionArgOfOp)
298 NameOS <<
"_r" << c.getArgIdx();
299 else if (c.getLoopDepth() >= 1)
300 NameOS <<
"_d" << c.getLoopDepth();
302 NameOS <<
"_s" << c.getOpPos();
305 return NameOS.str().str();
308void OpenMPDialect::initialize() {
311#include "mlir/Dialect/OpenMP/OpenMPOps.cpp.inc"
314#define GET_ATTRDEF_LIST
315#include "mlir/Dialect/OpenMP/OpenMPOpsAttributes.cpp.inc"
318#define GET_TYPEDEF_LIST
319#include "mlir/Dialect/OpenMP/OpenMPOpsTypes.cpp.inc"
322 declarePromisedInterface<ConvertToLLVMPatternInterface, OpenMPDialect>();
324 MemRefType::attachInterface<MemRefPointerLikeModel>(*
getContext());
325 LLVM::LLVMPointerType::attachInterface<LLVMPointerPointerLikeModel>(
330 mlir::ModuleOp::attachInterface<mlir::omp::OffloadModuleDefaultModel>(
336 mlir::LLVM::GlobalOp::attachInterface<
339 mlir::LLVM::LLVMFuncOp::attachInterface<
342 mlir::func::FuncOp::attachInterface<
351 if (!isa<DeclareTargetInterface>(op))
352 return op->
emitError() <<
"omp.declare_target can only be applied to "
353 "DeclareTargetInterface ops";
355 auto declareTargetAttr = dyn_cast<DeclareTargetAttr>(attr);
356 if (!declareTargetAttr)
358 <<
"omp.declare_target must be an #omp.declaretarget attribute";
360 if (isa<mlir::FunctionOpInterface>(op)) {
361 if (declareTargetAttr.getAutomap())
363 <<
"omp.declare_target 'automap' is not valid on functions";
366 if (declareTargetAttr.getCaptureClause() ==
367 mlir::omp::DeclareTargetCaptureClause::link)
369 <<
"omp.declare_target 'link' is not valid on functions";
372 if (declareTargetAttr.getImplicit())
374 <<
"omp.declare_target 'implicit' is only valid on functions";
380OpenMPDialect::verifyOperationAttribute(
Operation *op,
382 if (attribute.
getName() ==
"omp.declare_target")
410 allocatorVars.push_back(operand);
411 allocatorTypes.push_back(type);
417 allocateVars.push_back(operand);
418 allocateTypes.push_back(type);
429 for (
unsigned i = 0; i < allocateVars.size(); ++i) {
430 std::string separator = i == allocateVars.size() - 1 ?
"" :
", ";
431 p << allocatorVars[i] <<
" : " << allocatorTypes[i] <<
" -> ";
432 p << allocateVars[i] <<
" : " << allocateTypes[i] << separator;
440template <
typename ClauseAttr>
442 using ClauseT =
decltype(std::declval<ClauseAttr>().getValue());
447 if (std::optional<ClauseT> enumValue = symbolizeEnum<ClauseT>(enumStr)) {
448 attr = ClauseAttr::get(parser.
getContext(), *enumValue);
451 return parser.
emitError(loc,
"invalid clause value: '") << enumStr <<
"'";
454template <
typename ClauseAttr>
456 p << stringifyEnum(attr.getValue());
481 std::optional<omp::LinearModifier> linearModifier;
483 linearModifier = omp::LinearModifier::val;
485 linearModifier = omp::LinearModifier::ref;
487 linearModifier = omp::LinearModifier::uval;
490 bool hasLinearModifierParens = linearModifier.has_value();
491 if (hasLinearModifierParens && parser.
parseLParen())
499 if (hasLinearModifierParens && parser.
parseRParen())
502 linearVars.push_back(var);
503 linearTypes.push_back(type);
504 linearStepVars.push_back(stepVar);
505 linearStepTypes.push_back(stepType);
506 if (linearModifier) {
508 omp::LinearModifierAttr::get(parser.
getContext(), *linearModifier));
510 modifiers.push_back(UnitAttr::get(parser.
getContext()));
516 linearModifiers = ArrayAttr::get(parser.
getContext(), modifiers);
525 size_t linearVarsSize = linearVars.size();
526 for (
unsigned i = 0; i < linearVarsSize; ++i) {
530 Attribute modAttr = linearModifiers ? linearModifiers[i] :
nullptr;
531 auto mod = modAttr ? dyn_cast<omp::LinearModifierAttr>(modAttr) :
nullptr;
533 p << omp::stringifyLinearModifier(mod.getValue()) <<
"(";
535 p << linearVars[i] <<
" : " << linearTypes[i];
536 p <<
" = " << linearStepVars[i] <<
" : " << stepVarTypes[i];
552 if (!linearModifiers)
554 if (linearModifiers->size() != linearVars.size())
556 <<
"expected as many linear modifiers as linear variables";
557 if (!isDeclareSimd) {
558 for (
Attribute attr : *linearModifiers) {
561 auto modAttr = dyn_cast<omp::LinearModifierAttr>(attr);
564 omp::LinearModifier mod = modAttr.getValue();
565 if (mod == omp::LinearModifier::ref || mod == omp::LinearModifier::uval)
567 <<
"linear modifier '" << omp::stringifyLinearModifier(mod)
568 <<
"' may only be specified on a declare simd directive";
583 for (
const auto &it : nontemporalVars)
584 if (!nontemporalItems.insert(it).second)
585 return op->
emitOpError() <<
"nontemporal variable used more than once";
594 std::optional<ArrayAttr> alignments,
597 if (!alignedVars.empty()) {
598 if (!alignments || alignments->size() != alignedVars.size())
600 <<
"expected as many alignment values as aligned variables";
603 return op->
emitOpError() <<
"unexpected alignment values attribute";
609 for (
auto it : alignedVars)
610 if (!alignedItems.insert(it).second)
611 return op->
emitOpError() <<
"aligned variable used more than once";
617 for (
unsigned i = 0; i < (*alignments).size(); ++i) {
618 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>((*alignments)[i])) {
619 if (intAttr.getValue().sle(0))
620 return op->
emitOpError() <<
"alignment should be greater than 0";
622 return op->
emitOpError() <<
"expected integer alignment";
639 if (parser.parseOperand(alignedVars.emplace_back()) ||
640 parser.parseColonType(alignedTypes.emplace_back()) ||
641 parser.parseArrow() ||
642 parser.parseAttribute(alignmentVec.emplace_back())) {
649 alignmentsAttr = ArrayAttr::get(parser.getContext(), alignments);
656 std::optional<ArrayAttr> alignments) {
657 for (
unsigned i = 0; i < alignedVars.size(); ++i) {
660 p << alignedVars[i] <<
" : " << alignedVars[i].
getType();
661 p <<
" -> " << (*alignments)[i];
669 ArrayAttr privateSyms =
nullptr,
bool requirePrivateIndices =
false) {
670 if (allocateVars.size() != allocatorVars.size())
672 "expected equal sizes for allocate and allocator variables");
674 if (allocateVars.empty()) {
675 if (allocateAlignments)
677 "unexpected allocate alignments without allocate variables");
678 if (allocatePrivateIndices)
680 "unexpected allocate private indices without allocate variables");
684 if (allocateAlignments) {
686 if (alignments.size() != allocateVars.size())
688 "expected as many allocate alignments as allocate variables");
689 for (
int64_t alignment : alignments) {
691 return op->
emitError(
"expected non-negative allocate alignments");
692 if (alignment != 0 && (alignment & (alignment - 1)) != 0)
694 "expected positive allocate alignments to be powers of two");
698 if (!allocatePrivateIndices) {
699 if (requirePrivateIndices)
701 "expected an allocate private index for each allocate variable");
706 if (
indices.size() != allocateVars.size())
708 "expected as many allocate private indices as allocate variables");
711 for (
auto [allocateVar, privateIndex] :
712 llvm::zip_equal(allocateVars,
indices)) {
713 if (privateIndex < 0 ||
714 static_cast<uint64_t
>(privateIndex) >= privateVars.size())
715 return op->
emitError(
"allocate private index is out of range");
716 if (!usedPrivateSlots.insert(privateIndex).second)
718 "allocate private index refers to a private variable more than once");
720 Value privateVar = privateVars[privateIndex];
721 if (allocateVar.getType() != privateVar.
getType())
723 <<
"type mismatch between allocate variable and private variable "
726 if (allocateVar != privateVar)
728 <<
"allocate variable does not match private variable at index "
732 static_cast<uint64_t
>(privateIndex) >= privateSyms.size())
734 "allocate private index does not have a privatizer symbol");
736 auto privateSym = dyn_cast<SymbolRefAttr>(privateSyms[privateIndex]);
739 "allocate private index does not reference a privatizer symbol");
740 PrivateClauseOp privatizer =
743 return op->
emitError() <<
"failed to lookup privatizer op with symbol: '"
744 << privateSym <<
"'";
745 if (privatizer.getDataSharingType() != DataSharingClauseType::Private &&
746 privatizer.getDataSharingType() != DataSharingClauseType::FirstPrivate)
748 "allocate private index must refer to private or firstprivate "
762 if (modifiers.size() > 2)
764 for (
const auto &mod : modifiers) {
767 auto symbol = symbolizeScheduleModifier(mod);
770 <<
" unknown modifier type: " << mod;
775 if (modifiers.size() == 1) {
776 if (symbolizeScheduleModifier(modifiers[0]) == ScheduleModifier::simd) {
777 modifiers.push_back(modifiers[0]);
778 modifiers[0] = stringifyScheduleModifier(ScheduleModifier::none);
780 }
else if (modifiers.size() == 2) {
783 if (symbolizeScheduleModifier(modifiers[0]) == ScheduleModifier::simd ||
784 symbolizeScheduleModifier(modifiers[1]) != ScheduleModifier::simd)
786 <<
" incorrect modifier order";
802 ScheduleModifierAttr &scheduleMod, UnitAttr &scheduleSimd,
803 std::optional<OpAsmParser::UnresolvedOperand> &chunkSize,
808 std::optional<mlir::omp::ClauseScheduleKind> schedule =
809 symbolizeClauseScheduleKind(keyword);
813 scheduleAttr = ClauseScheduleKindAttr::get(parser.
getContext(), *schedule);
815 case ClauseScheduleKind::Static:
816 case ClauseScheduleKind::Dynamic:
817 case ClauseScheduleKind::Guided:
823 chunkSize = std::nullopt;
826 case ClauseScheduleKind::Auto:
827 case ClauseScheduleKind::Runtime:
828 case ClauseScheduleKind::Distribute:
829 chunkSize = std::nullopt;
838 modifiers.push_back(mod);
844 if (!modifiers.empty()) {
846 if (std::optional<ScheduleModifier> mod =
847 symbolizeScheduleModifier(modifiers[0])) {
848 scheduleMod = ScheduleModifierAttr::get(parser.
getContext(), *mod);
850 return parser.
emitError(loc,
"invalid schedule modifier");
853 if (modifiers.size() > 1) {
854 assert(symbolizeScheduleModifier(modifiers[1]) == ScheduleModifier::simd);
864 ClauseScheduleKindAttr scheduleKind,
865 ScheduleModifierAttr scheduleMod,
866 UnitAttr scheduleSimd,
Value scheduleChunk,
867 Type scheduleChunkType) {
868 p << stringifyClauseScheduleKind(scheduleKind.getValue());
870 p <<
" = " << scheduleChunk <<
" : " << scheduleChunk.
getType();
872 p <<
", " << stringifyScheduleModifier(scheduleMod.getValue());
884 ClauseOrderKindAttr &order,
885 OrderModifierAttr &orderMod) {
890 if (std::optional<OrderModifier> enumValue =
891 symbolizeOrderModifier(enumStr)) {
892 orderMod = OrderModifierAttr::get(parser.
getContext(), *enumValue);
899 if (std::optional<ClauseOrderKind> enumValue =
900 symbolizeClauseOrderKind(enumStr)) {
901 order = ClauseOrderKindAttr::get(parser.
getContext(), *enumValue);
904 return parser.
emitError(loc,
"invalid clause value: '") << enumStr <<
"'";
908 ClauseOrderKindAttr order,
909 OrderModifierAttr orderMod) {
911 p << stringifyOrderModifier(orderMod.getValue()) <<
":";
913 p << stringifyClauseOrderKind(order.getValue());
916template <
typename ClauseTypeAttr,
typename ClauseType>
919 std::optional<OpAsmParser::UnresolvedOperand> &operand,
921 std::optional<ClauseType> (*symbolizeClause)(StringRef),
922 StringRef clauseName) {
925 if (std::optional<ClauseType> enumValue = symbolizeClause(enumStr)) {
926 prescriptiveness = ClauseTypeAttr::get(parser.
getContext(), *enumValue);
931 <<
"invalid " << clauseName <<
" modifier : '" << enumStr <<
"'";
941 <<
"expected " << clauseName <<
" operand";
944 if (operand.has_value()) {
952template <
typename ClauseTypeAttr,
typename ClauseType>
955 ClauseTypeAttr prescriptiveness,
Value operand,
957 StringRef (*stringifyClauseType)(ClauseType)) {
959 if (prescriptiveness)
960 p << stringifyClauseType(prescriptiveness.getValue()) <<
", ";
963 p << operand <<
": " << operandType;
973 std::optional<OpAsmParser::UnresolvedOperand> &grainsize,
974 Type &grainsizeType) {
976 parser, grainsizeMod, grainsize, grainsizeType,
977 &symbolizeClauseGrainsizeType,
"grainsize");
981 ClauseGrainsizeTypeAttr grainsizeMod,
984 p, op, grainsizeMod, grainsize, grainsizeType,
985 &stringifyClauseGrainsizeType);
995 std::optional<OpAsmParser::UnresolvedOperand> &numTasks,
996 Type &numTasksType) {
998 parser, numTasksMod, numTasks, numTasksType, &symbolizeClauseNumTasksType,
1003 ClauseNumTasksTypeAttr numTasksMod,
1006 p, op, numTasksMod, numTasks, numTasksType, &stringifyClauseNumTasksType);
1022 return mlir::failure();
1023 inTypeAttr = TypeAttr::get(inType);
1052 if (!typeparams.empty()) {
1053 p <<
'(' << typeparams <<
" : " << typeparamsTypes <<
')';
1055 for (
auto sh :
shape) {
1067 FallbackModifierAttr fallback,
1068 Value dynGroupprivateSize) {
1069 if (!dynGroupprivateSize && (accessGroup || fallback))
1070 return op->
emitOpError(
"dyn_groupprivate modifiers require a size operand");
1076 OpAsmParser &parser, AccessGroupModifierAttr &accessGroupAttr,
1077 FallbackModifierAttr &fallbackAttr,
1078 std::optional<OpAsmParser::UnresolvedOperand> &dynGroupprivateSize,
1081 bool parsedAccessGroup =
false;
1082 bool parsedFallback =
false;
1083 bool parsedSize =
false;
1088 if (parsedAccessGroup)
1090 "duplicate access group modifier");
1091 accessGroupAttr = AccessGroupModifierAttr::get(
1092 parser.
getContext(), AccessGroupModifier::cgroup);
1093 parsedAccessGroup =
true;
1100 "duplicate fallback modifier");
1103 "expected '(' after 'fallback'");
1104 llvm::StringRef fbKind;
1108 "expected fallback modifier (abort/null/default_mem)");
1109 std::optional<FallbackModifier> fbEnum;
1110 if (fbKind ==
"abort")
1111 fbEnum = FallbackModifier::abort;
1112 else if (fbKind ==
"null")
1113 fbEnum = FallbackModifier::null;
1114 else if (fbKind ==
"default_mem")
1115 fbEnum = FallbackModifier::default_mem;
1118 "invalid fallback modifier '" + fbKind +
"'");
1119 fallbackAttr = FallbackModifierAttr::get(parser.
getContext(), *fbEnum);
1122 "expected ')' after fallback modifier");
1123 parsedFallback =
true;
1131 "duplicate size operand");
1132 dynGroupprivateSize = operand;
1136 "expected ':' and type after size operand");
1140 "expected dyn_groupprivate_size operand");
1145 AccessGroupModifierAttr modifierFirst,
1146 FallbackModifierAttr modifierSecond,
1147 Value dynGroupprivateSize,
1150 bool needsComma =
false;
1152 if (modifierFirst) {
1153 printer << modifierFirst.getValue();
1157 if (modifierSecond) {
1160 printer <<
"fallback(";
1161 printer << modifierSecond.getValue();
1166 if (dynGroupprivateSize) {
1169 printer << dynGroupprivateSize <<
" : " << sizeType;
1189 isByRefVec.push_back(parser.parseOptionalKeyword(
"byref").succeeded());
1190 if (parser.parseAttribute(symbolVec.emplace_back()) ||
1191 parser.parseOperand(inReductionVars.emplace_back()))
1201 [&]() { return parser.parseType(inReductionTypes.emplace_back()); }))
1204 if (inReductionVars.size() != inReductionTypes.size())
1209 inReductionSyms = ArrayAttr::get(parser.
getContext(), symbolAttrs);
1226 syms = ArrayAttr::get(ctx, values);
1235 llvm::interleaveComma(
1236 llvm::zip_equal(inReductionVars, syms.getValue(), byref.
asArrayRef()), p,
1238 auto [var, sym, isByRef] = t;
1246 llvm::interleaveComma(inReductionTypes, p);
1254struct MapParseArgs {
1255 SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars;
1256 SmallVectorImpl<Type> &types;
1257 MapParseArgs(SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars,
1258 SmallVectorImpl<Type> &types)
1259 : vars(vars), types(types) {}
1261struct PrivateParseArgs {
1262 llvm::SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars;
1263 llvm::SmallVectorImpl<Type> &types;
1265 UnitAttr &needsBarrier;
1267 PrivateParseArgs(SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars,
1268 SmallVectorImpl<Type> &types,
ArrayAttr &syms,
1269 UnitAttr &needsBarrier,
1271 : vars(vars), types(types), syms(syms), needsBarrier(needsBarrier),
1272 mapIndices(mapIndices) {}
1275struct ReductionParseArgs {
1276 SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars;
1277 SmallVectorImpl<Type> &types;
1280 ReductionModifierAttr *modifier;
1281 ReductionParseArgs(SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars,
1283 ArrayAttr &syms, ReductionModifierAttr *mod =
nullptr)
1284 : vars(vars), types(types), byref(byref), syms(syms), modifier(mod) {}
1287struct AllRegionParseArgs {
1288 std::optional<MapParseArgs> hasDeviceAddrArgs;
1289 std::optional<MapParseArgs> hostEvalArgs;
1290 std::optional<ReductionParseArgs> inReductionArgs;
1291 std::optional<MapParseArgs> mapArgs;
1292 std::optional<PrivateParseArgs> privateArgs;
1293 std::optional<ReductionParseArgs> reductionArgs;
1294 std::optional<ReductionParseArgs> taskReductionArgs;
1295 std::optional<MapParseArgs> useDeviceAddrArgs;
1296 std::optional<MapParseArgs> useDevicePtrArgs;
1301 return "private_barrier";
1311 ReductionModifierAttr *modifier =
nullptr,
1312 UnitAttr *needsBarrier =
nullptr) {
1316 unsigned regionArgOffset = regionPrivateArgs.size();
1326 std::optional<ReductionModifier> enumValue =
1327 symbolizeReductionModifier(enumStr);
1328 if (!enumValue.has_value())
1330 *modifier = ReductionModifierAttr::get(parser.
getContext(), *enumValue);
1337 isByRefVec.push_back(
1338 parser.parseOptionalKeyword(
"byref").succeeded());
1340 if (symbols && parser.parseAttribute(symbolVec.emplace_back()))
1343 if (parser.parseOperand(operands.emplace_back()) ||
1344 parser.parseArrow() ||
1345 parser.parseArgument(regionPrivateArgs.emplace_back()))
1349 if (parser.parseOptionalLSquare().succeeded()) {
1350 if (parser.parseKeyword(
"map_idx") || parser.parseEqual() ||
1351 parser.parseInteger(mapIndicesVec.emplace_back()) ||
1352 parser.parseRSquare())
1355 mapIndicesVec.push_back(-1);
1367 if (parser.parseType(types.emplace_back()))
1374 if (operands.size() != types.size())
1383 *needsBarrier = mlir::UnitAttr::get(parser.
getContext());
1386 auto *argsBegin = regionPrivateArgs.begin();
1388 argsBegin + regionArgOffset + types.size());
1389 for (
auto [prv, type] : llvm::zip_equal(argsSubrange, types)) {
1395 *symbols = ArrayAttr::get(parser.
getContext(), symbolAttrs);
1398 if (!mapIndicesVec.empty())
1411 StringRef keyword, std::optional<MapParseArgs> mapArgs) {
1426 StringRef keyword, std::optional<PrivateParseArgs> privateArgs) {
1432 parser, privateArgs->vars, privateArgs->types, entryBlockArgs,
1433 &privateArgs->syms, privateArgs->mapIndices,
nullptr,
1434 nullptr, &privateArgs->needsBarrier)))
1443 StringRef keyword, std::optional<ReductionParseArgs> reductionArgs) {
1448 parser, reductionArgs->vars, reductionArgs->types, entryBlockArgs,
1449 &reductionArgs->syms,
nullptr, &reductionArgs->byref,
1450 reductionArgs->modifier)))
1457 AllRegionParseArgs args) {
1461 args.hasDeviceAddrArgs)))
1463 <<
"invalid `has_device_addr` format";
1466 args.hostEvalArgs)))
1468 <<
"invalid `host_eval` format";
1471 args.inReductionArgs)))
1473 <<
"invalid `in_reduction` format";
1478 <<
"invalid `map_entries` format";
1483 <<
"invalid `private` format";
1486 args.reductionArgs)))
1488 <<
"invalid `reduction` format";
1491 args.taskReductionArgs)))
1493 <<
"invalid `task_reduction` format";
1496 args.useDeviceAddrArgs)))
1498 <<
"invalid `use_device_addr` format";
1501 args.useDevicePtrArgs)))
1503 <<
"invalid `use_device_addr` format";
1505 return parser.
parseRegion(region, entryBlockArgs);
1521 AllRegionParseArgs args;
1522 args.hasDeviceAddrArgs.emplace(hasDeviceAddrVars, hasDeviceAddrTypes);
1523 args.hostEvalArgs.emplace(hostEvalVars, hostEvalTypes);
1524 args.mapArgs.emplace(mapVars, mapTypes);
1525 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1526 privateNeedsBarrier, &privateMaps);
1537 UnitAttr &privateNeedsBarrier) {
1538 AllRegionParseArgs args;
1539 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1540 inReductionByref, inReductionSyms);
1541 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1542 privateNeedsBarrier);
1553 UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod,
1557 AllRegionParseArgs args;
1558 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1559 inReductionByref, inReductionSyms);
1560 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1561 privateNeedsBarrier);
1562 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1563 reductionSyms, &reductionMod);
1571 UnitAttr &privateNeedsBarrier) {
1572 AllRegionParseArgs args;
1573 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1574 privateNeedsBarrier);
1582 UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod,
1586 AllRegionParseArgs args;
1587 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1588 privateNeedsBarrier);
1589 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1590 reductionSyms, &reductionMod);
1599 AllRegionParseArgs args;
1600 args.taskReductionArgs.emplace(taskReductionVars, taskReductionTypes,
1601 taskReductionByref, taskReductionSyms);
1611 AllRegionParseArgs args;
1612 args.useDeviceAddrArgs.emplace(useDeviceAddrVars, useDeviceAddrTypes);
1613 args.useDevicePtrArgs.emplace(useDevicePtrVars, useDevicePtrTypes);
1622struct MapPrintArgs {
1627struct PrivatePrintArgs {
1631 UnitAttr needsBarrier;
1635 : vars(vars), types(types), syms(syms), needsBarrier(needsBarrier),
1636 mapIndices(mapIndices) {}
1638struct ReductionPrintArgs {
1643 ReductionModifierAttr modifier;
1645 ArrayAttr syms, ReductionModifierAttr mod =
nullptr)
1646 : vars(vars), types(types), byref(byref), syms(syms), modifier(mod) {}
1648struct AllRegionPrintArgs {
1649 std::optional<MapPrintArgs> hasDeviceAddrArgs;
1650 std::optional<MapPrintArgs> hostEvalArgs;
1651 std::optional<ReductionPrintArgs> inReductionArgs;
1652 std::optional<MapPrintArgs> mapArgs;
1653 std::optional<PrivatePrintArgs> privateArgs;
1654 std::optional<ReductionPrintArgs> reductionArgs;
1655 std::optional<ReductionPrintArgs> taskReductionArgs;
1656 std::optional<MapPrintArgs> useDeviceAddrArgs;
1657 std::optional<MapPrintArgs> useDevicePtrArgs;
1666 ReductionModifierAttr modifier =
nullptr, UnitAttr needsBarrier =
nullptr) {
1667 if (argsSubrange.empty())
1670 p << clauseName <<
"(";
1673 p <<
"mod: " << stringifyReductionModifier(modifier.getValue()) <<
", ";
1677 symbols = ArrayAttr::get(ctx, values);
1690 llvm::interleaveComma(llvm::zip_equal(operands, argsSubrange, symbols,
1691 mapIndices.asArrayRef(),
1692 byref.asArrayRef()),
1694 auto [op, arg, sym, map, isByRef] = t;
1700 p << op <<
" -> " << arg;
1703 p <<
" [map_idx=" << map <<
"]";
1706 llvm::interleaveComma(types, p);
1714 StringRef clauseName,
ValueRange argsSubrange,
1715 std::optional<MapPrintArgs> mapArgs) {
1722 StringRef clauseName,
ValueRange argsSubrange,
1723 std::optional<PrivatePrintArgs> privateArgs) {
1726 p, ctx, clauseName, argsSubrange, privateArgs->vars, privateArgs->types,
1727 privateArgs->syms, privateArgs->mapIndices,
nullptr,
1728 nullptr, privateArgs->needsBarrier);
1734 std::optional<ReductionPrintArgs> reductionArgs) {
1737 reductionArgs->vars, reductionArgs->types,
1738 reductionArgs->syms,
nullptr,
1739 reductionArgs->byref, reductionArgs->modifier);
1743 const AllRegionPrintArgs &args) {
1744 auto iface = llvm::cast<mlir::omp::BlockArgOpenMPOpInterface>(op);
1748 iface.getHasDeviceAddrBlockArgs(),
1749 args.hasDeviceAddrArgs);
1753 args.inReductionArgs);
1759 args.reductionArgs);
1761 iface.getTaskReductionBlockArgs(),
1762 args.taskReductionArgs);
1764 iface.getUseDeviceAddrBlockArgs(),
1765 args.useDeviceAddrArgs);
1767 iface.getUseDevicePtrBlockArgs(), args.useDevicePtrArgs);
1781 UnitAttr privateNeedsBarrier,
1783 AllRegionPrintArgs args;
1784 args.hasDeviceAddrArgs.emplace(hasDeviceAddrVars, hasDeviceAddrTypes);
1785 args.hostEvalArgs.emplace(hostEvalVars, hostEvalTypes);
1786 args.mapArgs.emplace(mapVars, mapTypes);
1787 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1788 privateNeedsBarrier, privateMaps);
1796 ArrayAttr privateSyms, UnitAttr privateNeedsBarrier) {
1797 AllRegionPrintArgs args;
1798 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1799 inReductionByref, inReductionSyms);
1800 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1801 privateNeedsBarrier,
1810 ArrayAttr privateSyms, UnitAttr privateNeedsBarrier,
1811 ReductionModifierAttr reductionMod,
ValueRange reductionVars,
1814 AllRegionPrintArgs args;
1815 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1816 inReductionByref, inReductionSyms);
1817 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1818 privateNeedsBarrier,
1820 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1821 reductionSyms, reductionMod);
1828 UnitAttr privateNeedsBarrier) {
1829 AllRegionPrintArgs args;
1830 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1831 privateNeedsBarrier,
1839 ReductionModifierAttr reductionMod,
ValueRange reductionVars,
1842 AllRegionPrintArgs args;
1843 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1844 privateNeedsBarrier,
1846 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1847 reductionSyms, reductionMod);
1857 AllRegionPrintArgs args;
1858 args.taskReductionArgs.emplace(taskReductionVars, taskReductionTypes,
1859 taskReductionByref, taskReductionSyms);
1869 AllRegionPrintArgs args;
1870 args.useDeviceAddrArgs.emplace(useDeviceAddrVars, useDeviceAddrTypes);
1871 args.useDevicePtrArgs.emplace(useDevicePtrVars, useDevicePtrTypes);
1875template <
typename ParsePrefixFn>
1884 if (failed(parsePrefix()))
1892 if (llvm::isa<mlir::omp::IteratedType>(ty)) {
1893 iteratedVars.push_back(v);
1894 iteratedTypes.push_back(ty);
1896 plainVars.push_back(v);
1897 plainTypes.push_back(ty);
1903template <
typename Pr
intPrefixFn>
1907 PrintPrefixFn &&printPrefixForPlain,
1908 PrintPrefixFn &&printPrefixForIterated) {
1915 p << v <<
" : " << t;
1919 for (
unsigned i = 0; i < iteratedVars.size(); ++i)
1920 emit(iteratedVars[i], iteratedTypes[i], printPrefixForIterated);
1921 for (
unsigned i = 0; i < plainVars.size(); ++i)
1922 emit(plainVars[i], plainTypes[i], printPrefixForPlain);
1930 if (!reductionVars.empty()) {
1931 if (!reductionSyms || reductionSyms->size() != reductionVars.size())
1933 <<
"expected as many reduction symbol references "
1934 "as reduction variables";
1935 if (reductionByref && reductionByref->size() != reductionVars.size())
1936 return op->
emitError() <<
"expected as many reduction variable by "
1937 "reference attributes as reduction variables";
1940 return op->
emitOpError() <<
"unexpected reduction symbol references";
1947 for (
auto args : llvm::zip(reductionVars, *reductionSyms)) {
1948 Value accum = std::get<0>(args);
1950 if (!accumulators.insert(accum).second)
1951 return op->
emitOpError() <<
"accumulator variable used more than once";
1954 auto symbolRef = llvm::cast<SymbolRefAttr>(std::get<1>(args));
1958 return op->
emitOpError() <<
"expected symbol reference " << symbolRef
1959 <<
" to point to a reduction declaration";
1961 if (decl.getAccumulatorType() && decl.getAccumulatorType() != varType)
1963 <<
"expected accumulator (" << varType
1964 <<
") to be the same type as reduction declaration ("
1965 << decl.getAccumulatorType() <<
")";
1984 if (parser.parseOperand(copyprivateVars.emplace_back()) ||
1985 parser.parseArrow() ||
1986 parser.parseAttribute(symsVec.emplace_back()) ||
1987 parser.parseColonType(copyprivateTypes.emplace_back()))
1993 copyprivateSyms = ArrayAttr::get(parser.
getContext(), syms);
2001 std::optional<ArrayAttr> copyprivateSyms) {
2002 if (!copyprivateSyms.has_value())
2004 llvm::interleaveComma(
2005 llvm::zip(copyprivateVars, *copyprivateSyms, copyprivateTypes), p,
2006 [&](
const auto &args) {
2007 p << std::get<0>(args) <<
" -> " << std::get<1>(args) <<
" : "
2008 << std::get<2>(args);
2015 std::optional<ArrayAttr> copyprivateSyms) {
2016 size_t copyprivateSymsSize =
2017 copyprivateSyms.has_value() ? copyprivateSyms->size() : 0;
2018 if (copyprivateSymsSize != copyprivateVars.size())
2019 return op->
emitOpError() <<
"inconsistent number of copyprivate vars (= "
2020 << copyprivateVars.size()
2021 <<
") and functions (= " << copyprivateSymsSize
2022 <<
"), both must be equal";
2023 if (!copyprivateSyms.has_value())
2026 for (
auto copyprivateVarAndSym :
2027 llvm::zip(copyprivateVars, *copyprivateSyms)) {
2029 llvm::cast<SymbolRefAttr>(std::get<1>(copyprivateVarAndSym));
2030 std::optional<std::variant<mlir::func::FuncOp, mlir::LLVM::LLVMFuncOp>>
2032 if (mlir::func::FuncOp mlirFuncOp =
2035 funcOp = mlirFuncOp;
2036 else if (mlir::LLVM::LLVMFuncOp llvmFuncOp =
2039 funcOp = llvmFuncOp;
2041 auto getNumArguments = [&] {
2042 return std::visit([](
auto &f) {
return f.getNumArguments(); }, *funcOp);
2045 auto getArgumentType = [&](
unsigned i) {
2046 return std::visit([i](
auto &f) {
return f.getArgumentTypes()[i]; },
2051 return op->
emitOpError() <<
"expected symbol reference " << symbolRef
2052 <<
" to point to a copy function";
2054 if (getNumArguments() != 2)
2056 <<
"expected copy function " << symbolRef <<
" to have 2 operands";
2058 Type argTy = getArgumentType(0);
2059 if (argTy != getArgumentType(1))
2060 return op->
emitOpError() <<
"expected copy function " << symbolRef
2061 <<
" arguments to have the same type";
2063 Type varType = std::get<0>(copyprivateVarAndSym).getType();
2064 if (argTy != varType)
2066 <<
"expected copy function arguments' type (" << argTy
2067 <<
") to be the same as copyprivate variable's type (" << varType
2092 OpAsmParser::UnresolvedOperand operand;
2094 if (parser.parseKeyword(&keyword) || parser.parseArrow() ||
2095 parser.parseOperand(operand) || parser.parseColonType(ty))
2097 std::optional<ClauseTaskDepend> keywordDepend =
2098 symbolizeClauseTaskDepend(keyword);
2102 ClauseTaskDependAttr::get(parser.getContext(), *keywordDepend);
2103 if (llvm::isa<mlir::omp::IteratedType>(ty)) {
2104 iteratedVars.push_back(operand);
2105 iteratedTypes.push_back(ty);
2106 iterKindsVec.push_back(kindAttr);
2108 dependVars.push_back(operand);
2109 dependTypes.push_back(ty);
2110 kindsVec.push_back(kindAttr);
2116 dependKinds = ArrayAttr::get(parser.
getContext(), kinds);
2118 iteratedKinds = ArrayAttr::get(parser.
getContext(), iterKinds);
2125 std::optional<ArrayAttr> dependKinds,
2128 std::optional<ArrayAttr> iteratedKinds) {
2131 std::optional<ArrayAttr> kinds) {
2132 for (
unsigned i = 0, e = vars.size(); i < e; ++i) {
2135 p << stringifyClauseTaskDepend(
2136 llvm::cast<mlir::omp::ClauseTaskDependAttr>((*kinds)[i])
2138 <<
" -> " << vars[i] <<
" : " << types[i];
2142 printEntries(dependVars, dependTypes, dependKinds);
2143 printEntries(iteratedVars, iteratedTypes, iteratedKinds);
2148 std::optional<ArrayAttr> dependKinds,
2150 std::optional<ArrayAttr> iteratedKinds,
2152 if (!dependVars.empty()) {
2153 if (!dependKinds || dependKinds->size() != dependVars.size())
2154 return op->
emitOpError() <<
"expected as many depend values"
2155 " as depend variables";
2157 if (dependKinds && !dependKinds->empty())
2158 return op->
emitOpError() <<
"unexpected depend values";
2161 if (!iteratedVars.empty()) {
2162 if (!iteratedKinds || iteratedKinds->size() != iteratedVars.size())
2163 return op->
emitOpError() <<
"expected as many depend iterated values"
2164 " as depend iterated variables";
2166 if (iteratedKinds && !iteratedKinds->empty())
2167 return op->
emitOpError() <<
"unexpected depend iterated values";
2182 IntegerAttr &hintAttr) {
2183 StringRef hintKeyword;
2189 auto parseKeyword = [&]() -> ParseResult {
2192 if (hintKeyword ==
"uncontended")
2194 else if (hintKeyword ==
"contended")
2196 else if (hintKeyword ==
"nonspeculative")
2198 else if (hintKeyword ==
"speculative")
2202 << hintKeyword <<
" is not a valid hint";
2213 IntegerAttr hintAttr) {
2214 int64_t hint = hintAttr.getInt();
2222 auto bitn = [](
int value,
int n) ->
bool {
return value & (1 << n); };
2224 bool uncontended = bitn(hint, 0);
2225 bool contended = bitn(hint, 1);
2226 bool nonspeculative = bitn(hint, 2);
2227 bool speculative = bitn(hint, 3);
2231 hints.push_back(
"uncontended");
2233 hints.push_back(
"contended");
2235 hints.push_back(
"nonspeculative");
2237 hints.push_back(
"speculative");
2239 llvm::interleaveComma(hints, p);
2246 auto bitn = [](
int value,
int n) ->
bool {
return value & (1 << n); };
2248 bool uncontended = bitn(hint, 0);
2249 bool contended = bitn(hint, 1);
2250 bool nonspeculative = bitn(hint, 2);
2251 bool speculative = bitn(hint, 3);
2253 if (uncontended && contended)
2254 return op->
emitOpError() <<
"the hints omp_sync_hint_uncontended and "
2255 "omp_sync_hint_contended cannot be combined";
2256 if (nonspeculative && speculative)
2257 return op->
emitOpError() <<
"the hints omp_sync_hint_nonspeculative and "
2258 "omp_sync_hint_speculative cannot be combined.";
2269 return (value & flag) == flag;
2277static ParseResult parseMapClause(
OpAsmParser &parser,
2278 ClauseMapFlagsAttr &mapType) {
2279 ClauseMapFlags mapTypeBits = ClauseMapFlags::none;
2282 auto parseTypeAndMod = [&]() -> ParseResult {
2283 StringRef mapTypeMod;
2287 if (mapTypeMod ==
"always")
2288 mapTypeBits |= ClauseMapFlags::always;
2290 if (mapTypeMod ==
"implicit")
2291 mapTypeBits |= ClauseMapFlags::implicit;
2293 if (mapTypeMod ==
"ompx_hold")
2294 mapTypeBits |= ClauseMapFlags::ompx_hold;
2296 if (mapTypeMod ==
"close")
2297 mapTypeBits |= ClauseMapFlags::close;
2299 if (mapTypeMod ==
"present")
2300 mapTypeBits |= ClauseMapFlags::present;
2302 if (mapTypeMod ==
"to")
2303 mapTypeBits |= ClauseMapFlags::to;
2305 if (mapTypeMod ==
"from")
2306 mapTypeBits |= ClauseMapFlags::from;
2308 if (mapTypeMod ==
"tofrom")
2309 mapTypeBits |= ClauseMapFlags::to | ClauseMapFlags::from;
2311 if (mapTypeMod ==
"delete")
2312 mapTypeBits |= ClauseMapFlags::del;
2314 if (mapTypeMod ==
"storage")
2315 mapTypeBits |= ClauseMapFlags::storage;
2317 if (mapTypeMod ==
"return_param")
2318 mapTypeBits |= ClauseMapFlags::return_param;
2320 if (mapTypeMod ==
"private")
2321 mapTypeBits |= ClauseMapFlags::priv;
2323 if (mapTypeMod ==
"literal")
2324 mapTypeBits |= ClauseMapFlags::literal;
2326 if (mapTypeMod ==
"attach")
2327 mapTypeBits |= ClauseMapFlags::attach;
2329 if (mapTypeMod ==
"attach_always")
2330 mapTypeBits |= ClauseMapFlags::attach_always;
2332 if (mapTypeMod ==
"attach_never")
2333 mapTypeBits |= ClauseMapFlags::attach_never;
2335 if (mapTypeMod ==
"attach_auto")
2336 mapTypeBits |= ClauseMapFlags::attach_auto;
2338 if (mapTypeMod ==
"ref_ptr")
2339 mapTypeBits |= ClauseMapFlags::ref_ptr;
2341 if (mapTypeMod ==
"ref_ptee")
2342 mapTypeBits |= ClauseMapFlags::ref_ptee;
2344 if (mapTypeMod ==
"is_device_ptr")
2345 mapTypeBits |= ClauseMapFlags::is_device_ptr;
2347 if (mapTypeMod ==
"target_param")
2348 mapTypeBits |= ClauseMapFlags::target_param;
2365 ClauseMapFlagsAttr mapType) {
2367 ClauseMapFlags mapFlags = mapType.getValue();
2372 mapTypeStrs.push_back(
"always");
2374 mapTypeStrs.push_back(
"implicit");
2376 mapTypeStrs.push_back(
"ompx_hold");
2378 mapTypeStrs.push_back(
"close");
2380 mapTypeStrs.push_back(
"present");
2382 mapTypeStrs.push_back(
"target_param");
2391 mapTypeStrs.push_back(
"tofrom");
2393 mapTypeStrs.push_back(
"from");
2395 mapTypeStrs.push_back(
"to");
2398 mapTypeStrs.push_back(
"delete");
2400 mapTypeStrs.push_back(
"return_param");
2402 mapTypeStrs.push_back(
"storage");
2404 mapTypeStrs.push_back(
"private");
2406 mapTypeStrs.push_back(
"literal");
2408 mapTypeStrs.push_back(
"attach");
2410 mapTypeStrs.push_back(
"attach_always");
2412 mapTypeStrs.push_back(
"attach_never");
2414 mapTypeStrs.push_back(
"attach_auto");
2416 mapTypeStrs.push_back(
"ref_ptr");
2418 mapTypeStrs.push_back(
"ref_ptee");
2420 mapTypeStrs.push_back(
"is_device_ptr");
2421 if (mapFlags == ClauseMapFlags::none)
2422 mapTypeStrs.push_back(
"none");
2424 for (
unsigned int i = 0; i < mapTypeStrs.size(); ++i) {
2425 p << mapTypeStrs[i];
2426 if (i + 1 < mapTypeStrs.size()) {
2432static ParseResult parseMembersIndex(
OpAsmParser &parser,
2436 auto parseIndices = [&]() -> ParseResult {
2441 APInt(64, value,
false)));
2455 memberIdxs.push_back(ArrayAttr::get(parser.
getContext(), values));
2459 if (!memberIdxs.empty())
2460 membersIdx = ArrayAttr::get(parser.
getContext(), memberIdxs);
2470 llvm::interleaveComma(membersIdx, p, [&p](
Attribute v) {
2472 auto memberIdx = cast<ArrayAttr>(v);
2473 llvm::interleaveComma(memberIdx.getValue(), p, [&p](
Attribute v2) {
2474 p << cast<IntegerAttr>(v2).getInt();
2481 VariableCaptureKindAttr mapCaptureType) {
2482 std::string typeCapStr;
2483 llvm::raw_string_ostream typeCap(typeCapStr);
2484 if (mapCaptureType.getValue() == mlir::omp::VariableCaptureKind::ByRef)
2486 if (mapCaptureType.getValue() == mlir::omp::VariableCaptureKind::ByCopy)
2487 typeCap <<
"ByCopy";
2488 if (mapCaptureType.getValue() == mlir::omp::VariableCaptureKind::VLAType)
2489 typeCap <<
"VLAType";
2490 if (mapCaptureType.getValue() == mlir::omp::VariableCaptureKind::This)
2496 VariableCaptureKindAttr &mapCaptureType) {
2497 StringRef mapCaptureKey;
2501 if (mapCaptureKey ==
"This")
2502 mapCaptureType = mlir::omp::VariableCaptureKindAttr::get(
2503 parser.
getContext(), mlir::omp::VariableCaptureKind::This);
2504 if (mapCaptureKey ==
"ByRef")
2505 mapCaptureType = mlir::omp::VariableCaptureKindAttr::get(
2506 parser.
getContext(), mlir::omp::VariableCaptureKind::ByRef);
2507 if (mapCaptureKey ==
"ByCopy")
2508 mapCaptureType = mlir::omp::VariableCaptureKindAttr::get(
2509 parser.
getContext(), mlir::omp::VariableCaptureKind::ByCopy);
2510 if (mapCaptureKey ==
"VLAType")
2511 mapCaptureType = mlir::omp::VariableCaptureKindAttr::get(
2512 parser.
getContext(), mlir::omp::VariableCaptureKind::VLAType);
2518 Operation *op, mlir::omp::MapInfoOp mapInfoOp,
2522 mlir::omp::ClauseMapFlags mapTypeBits = mapInfoOp.getMapType();
2525 bool from =
mapTypeToBool(mapTypeBits, ClauseMapFlags::from);
2528 bool always =
mapTypeToBool(mapTypeBits, ClauseMapFlags::always);
2529 bool close =
mapTypeToBool(mapTypeBits, ClauseMapFlags::close);
2530 bool implicit =
mapTypeToBool(mapTypeBits, ClauseMapFlags::implicit);
2531 bool attach =
mapTypeToBool(mapTypeBits, ClauseMapFlags::attach);
2533 if ((isa<TargetDataOp>(op) || isa<TargetOp>(op)) && del)
2535 "to, from, tofrom and alloc map types are permitted");
2537 if (isa<TargetEnterDataOp>(op) && (from || del))
2538 return emitError(op->
getLoc(),
"to and alloc map types are permitted");
2540 if (isa<TargetExitDataOp>(op) && to)
2542 "from, release and delete map types are permitted");
2544 if (isa<TargetUpdateOp>(op)) {
2547 "at least one of to or from map types must be "
2548 "specified, other map types are not permitted");
2551 if (!to && !from && !attach) {
2553 "at least one of to or from or attach map types must be "
2554 "specified, other map types are not permitted");
2557 auto updateVar = mapInfoOp.getVarPtr();
2559 if ((to && from) || (to && updateFromVars.contains(updateVar)) ||
2560 (from && updateToVars.contains(updateVar))) {
2563 "either to or from map types can be specified, not both");
2566 if (always || close || implicit) {
2569 "present, mapper and iterator map type modifiers are permitted");
2575 to ? updateToVars.insert(updateVar) : updateFromVars.insert(updateVar);
2579 if ((mapInfoOp.getVarPtrPtr() && !mapInfoOp.getVarPtrPtrType()) ||
2580 (!mapInfoOp.getVarPtrPtr() && mapInfoOp.getVarPtrPtrType())) {
2582 "if varPtrPtr or varPtrPtrType is specified, then both "
2594 for (
auto mapOp : mapVars) {
2595 if (!mapOp.getDefiningOp())
2598 if (
auto mapInfoOp = mapOp.getDefiningOp<mlir::omp::MapInfoOp>()) {
2602 }
else if (!isa<DeclareMapperInfoOp>(op)) {
2604 "map argument is not a map entry operation");
2609 for (
auto iterVal : mapIterated) {
2610 auto iterOp = iterVal.getDefiningOp<mlir::omp::IteratorOp>();
2612 return op->
emitOpError() <<
"'map_iterated' arguments must be defined by "
2613 "'omp.iterator' ops";
2617 cast<mlir::omp::YieldOp>(iterOp.getRegion().front().getTerminator());
2618 auto yieldedMapInfo =
2619 yieldOp.getResults()[0].getDefiningOp<mlir::omp::MapInfoOp>();
2620 if (!yieldedMapInfo)
2621 return op->
emitOpError() <<
"'map_iterated' iterator body must yield "
2622 "a value defined by 'omp.map.info'";
2632template <
typename OpType>
2636 std::optional<DenseI64ArrayAttr> privateMapIndices =
2637 targetOp.getPrivateMapsAttr();
2640 if (!privateMapIndices.has_value() || !privateMapIndices.value())
2645 if (privateMapIndices.value().size() !=
2646 static_cast<int64_t>(privateVars.size()))
2647 return emitError(targetOp.getLoc(),
"sizes of `private` operand range and "
2648 "`private_maps` attribute mismatch");
2658 StringRef clauseName,
2660 for (
Value var : vars)
2661 if (!llvm::isa_and_present<MapInfoOp>(var.getDefiningOp()))
2663 <<
"'" << clauseName
2664 <<
"' arguments must be defined by 'omp.map.info' ops";
2668LogicalResult MapInfoOp::verify() {
2669 if (getMapperId() &&
2671 *
this, getMapperIdAttr())) {
2686 const TargetDataOperands &clauses) {
2687 TargetDataOp::build(builder, state, clauses.device, clauses.ifExpr,
2688 clauses.mapVars, clauses.mapIterated,
2689 clauses.useDeviceAddrVars, clauses.useDevicePtrVars);
2692LogicalResult TargetDataOp::verify() {
2693 if (getMapVars().empty() && getMapIterated().empty() &&
2694 getUseDevicePtrVars().empty() && getUseDeviceAddrVars().empty()) {
2695 return ::emitError(this->getLoc(),
2696 "At least one of map, use_device_ptr_vars, or "
2697 "use_device_addr_vars operand must be present");
2701 getUseDevicePtrVars())))
2705 getUseDeviceAddrVars())))
2715void TargetEnterDataOp::build(
2719 TargetEnterDataOp::build(
2721 clauses.dependVars,
makeArrayAttr(ctx, clauses.dependIteratedKinds),
2722 clauses.dependIterated, clauses.device, clauses.ifExpr, clauses.mapVars,
2723 clauses.mapIterated, clauses.nowait);
2726LogicalResult TargetEnterDataOp::verify() {
2727 LogicalResult verifyDependVars =
2729 getDependIteratedKinds(), getDependIterated());
2730 return failed(verifyDependVars)
2742 TargetExitDataOp::build(
2744 clauses.dependVars,
makeArrayAttr(ctx, clauses.dependIteratedKinds),
2745 clauses.dependIterated, clauses.device, clauses.ifExpr, clauses.mapVars,
2746 clauses.mapIterated, clauses.nowait);
2749LogicalResult TargetExitDataOp::verify() {
2750 LogicalResult verifyDependVars =
2752 getDependIteratedKinds(), getDependIterated());
2753 return failed(verifyDependVars)
2765 TargetUpdateOp::build(builder, state,
makeArrayAttr(ctx, clauses.dependKinds),
2768 clauses.dependIterated, clauses.device, clauses.ifExpr,
2769 clauses.mapVars, clauses.mapIterated, clauses.nowait);
2772LogicalResult TargetUpdateOp::verify() {
2773 LogicalResult verifyDependVars =
2775 getDependIteratedKinds(), getDependIterated());
2776 return failed(verifyDependVars)
2789 builder, state, clauses.allocateVars, clauses.allocatorVars,
2792 makeArrayAttr(ctx, clauses.dependKinds), clauses.dependVars,
2793 makeArrayAttr(ctx, clauses.dependIteratedKinds), clauses.dependIterated,
2794 clauses.device, clauses.dynGroupprivateAccessGroup,
2795 clauses.dynGroupprivateFallback, clauses.dynGroupprivateSize,
2796 clauses.hasDeviceAddrVars, clauses.hostEvalVars, clauses.ifExpr,
2797 clauses.inReductionVars,
2799 makeArrayAttr(ctx, clauses.inReductionSyms), clauses.isDevicePtrVars,
2800 clauses.mapVars, clauses.mapIterated, clauses.nowait, clauses.privateVars,
2801 makeArrayAttr(ctx, clauses.privateSyms), clauses.privateNeedsBarrier,
2802 clauses.threadLimitVars,
nullptr, clauses.
kernelType);
2805bool TargetOp::hasHostEvalTripCount() {
2806 TargetExecMode mode = getKernelType();
2807 if (mode == TargetExecMode::spmd || mode == TargetExecMode::spmd_no_loop)
2810 if (mode == TargetExecMode::bare)
2816 cast<ComposableOpInterface>(getOperation()).findCapturedOp();
2817 if (
auto loopNestOp = dyn_cast_if_present<LoopNestOp>(capturedOp)) {
2819 loopNestOp.gatherWrappers(loopWrappers);
2821 LoopWrapperInterface *innermostWrapper = loopWrappers.begin();
2822 if (isa<SimdOp>(innermostWrapper))
2823 innermostWrapper = std::next(innermostWrapper);
2825 auto numWrappers = std::distance(innermostWrapper, loopWrappers.end());
2826 if (numWrappers != 1)
2829 if (!isa<DistributeOp>(innermostWrapper))
2833 if (isa_and_present<TeamsOp>(parentOp) &&
2849 if (mapVarPtr == inReductionVar)
2855LogicalResult TargetOp::verify() {
2857 getOperation(), getAllocateVars(), getAllocatorVars(),
2858 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
2859 getPrivateVars(), getPrivateSymsAttr())))
2862 if (getKernelType() == TargetExecMode::bare && !isCombined())
2863 return emitOpError() <<
"bare kernel requires 'omp.combined'";
2866 getDependIteratedKinds(),
2867 getDependIterated())))
2871 getHasDeviceAddrVars())))
2878 *
this, getDynGroupprivateAccessGroupAttr(),
2879 getDynGroupprivateFallbackAttr(), getDynGroupprivateSize())))
2886 getInReductionVars(),
2887 getInReductionByref())))
2895 for (
Value inReductionVar : getInReductionVars()) {
2896 bool captured =
false;
2897 for (
Value mapVar : getMapVars()) {
2898 auto mapInfo = mapVar.getDefiningOp<MapInfoOp>();
2905 return emitOpError() <<
"in_reduction variable must be captured by a "
2906 "matching map_entries entry";
2912LogicalResult TargetOp::verifyRegions() {
2913 auto teamsOps = getOps<TeamsOp>();
2914 auto numNestedTeams = std::distance(teamsOps.begin(), teamsOps.end());
2915 if (numNestedTeams > 1)
2916 return emitError(
"target containing multiple 'omp.teams' nested ops");
2918 if (numNestedTeams == 0) {
2919 switch (getKernelType()) {
2920 case TargetExecMode::bare:
2921 return emitOpError()
2922 <<
"bare kernel must contain a nested 'omp.teams' operation";
2923 case TargetExecMode::spmd_no_loop:
2924 return emitOpError() <<
"spmd_no_loop kernel must contain a nested "
2925 "'omp.teams' operation";
2932 cast<ComposableOpInterface>(getOperation()).findCapturedOp();
2933 if ((getKernelType() == TargetExecMode::spmd ||
2934 getKernelType() == TargetExecMode::spmd_no_loop) &&
2935 !isa_and_present<LoopNestOp>(capturedOp))
2936 return emitOpError()
2937 <<
"SPMD kernel must capture an 'omp.loop_nest' operation";
2939 bool isTargetDevice =
false;
2940 if (
auto offloadMod = (*this)->getParentOfType<OffloadModuleInterface>())
2941 if (offloadMod.getIsTargetDevice())
2942 isTargetDevice =
true;
2946 cast<BlockArgOpenMPOpInterface>(getOperation()).getHostEvalBlockArgs();
2948 bool hostEvalTripCount = hasHostEvalTripCount();
2949 for (
Value hostEvalArg : hostEvalBlockArgs) {
2951 if (
auto teamsOp = dyn_cast<TeamsOp>(user)) {
2953 if (hostEvalArg == teamsOp.getNumTeamsLower() ||
2954 llvm::is_contained(teamsOp.getNumTeamsUpperVars(), hostEvalArg) ||
2955 llvm::is_contained(teamsOp.getThreadLimitVars(), hostEvalArg))
2958 return emitOpError() <<
"host_eval argument only legal as 'num_teams' "
2959 "and 'thread_limit' in 'omp.teams'";
2961 if (
auto parallelOp = dyn_cast<ParallelOp>(user)) {
2962 if (llvm::is_contained(parallelOp.getNumThreadsVars(), hostEvalArg))
2965 return emitOpError()
2966 <<
"host_eval argument only legal as 'num_threads' in "
2969 if (
auto loopNestOp = dyn_cast<LoopNestOp>(user)) {
2970 if (hostEvalTripCount &&
2971 (llvm::is_contained(loopNestOp.getLoopLowerBounds(), hostEvalArg) ||
2972 llvm::is_contained(loopNestOp.getLoopUpperBounds(), hostEvalArg) ||
2973 llvm::is_contained(loopNestOp.getLoopSteps(), hostEvalArg)))
2976 return emitOpError() <<
"host_eval argument only legal as loop bounds "
2977 "and steps in 'omp.loop_nest' when trip count "
2978 "must be evaluated in the host";
2981 return emitOpError() <<
"host_eval argument illegal use in '"
2982 << user->getName() <<
"' operation";
2986 if (hostEvalTripCount && !isTargetDevice) {
2987 auto loopOp = cast<LoopNestOp>(capturedOp);
2988 for (
auto arg : llvm::concat<Value>(loopOp.getLoopLowerBounds(),
2989 loopOp.getLoopUpperBounds(),
2990 loopOp.getLoopSteps())) {
2991 if (!llvm::is_contained(hostEvalBlockArgs, arg))
2992 return emitOpError() <<
"nested 'omp.loop_nest' bounds expected to "
2993 "be host-evaluated";
3006 ParallelOp::build(builder, state,
ValueRange(),
3020 const ParallelOperands &clauses) {
3022 ParallelOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3025 clauses.ifExpr, clauses.numThreadsVars, clauses.privateVars,
3027 clauses.privateNeedsBarrier, clauses.procBindKind,
3028 clauses.reductionMod, clauses.reductionVars,
3033template <
typename OpType>
3035 auto privateVars = op.getPrivateVars();
3036 auto privateSyms = op.getPrivateSymsAttr();
3038 if (privateVars.empty() && (privateSyms ==
nullptr || privateSyms.empty()))
3041 auto numPrivateVars = privateVars.size();
3042 auto numPrivateSyms = (privateSyms ==
nullptr) ? 0 : privateSyms.size();
3044 if (numPrivateVars != numPrivateSyms)
3045 return op.emitError() <<
"inconsistent number of private variables and "
3046 "privatizer op symbols, private vars: "
3048 <<
" vs. privatizer op symbols: " << numPrivateSyms;
3050 for (
auto privateVarInfo : llvm::zip_equal(privateVars, privateSyms)) {
3051 Type varType = std::get<0>(privateVarInfo).getType();
3052 SymbolRefAttr privateSym = cast<SymbolRefAttr>(std::get<1>(privateVarInfo));
3053 PrivateClauseOp privatizerOp =
3056 if (privatizerOp ==
nullptr)
3057 return op.emitError() <<
"failed to lookup privatizer op with symbol: '"
3058 << privateSym <<
"'";
3060 Type privatizerType = privatizerOp.getArgType();
3062 if (privatizerType && (varType != privatizerType))
3063 return op.emitError()
3064 <<
"type mismatch between a "
3065 << (privatizerOp.getDataSharingType() ==
3066 DataSharingClauseType::Private
3069 <<
" variable and its privatizer op, var type: " << varType
3070 <<
" vs. privatizer op type: " << privatizerType;
3076LogicalResult ParallelOp::verify() {
3080 getOperation(), getAllocateVars(), getAllocatorVars(),
3081 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3082 getPrivateVars(), getPrivateSymsAttr(),
3087 getReductionByref());
3090LogicalResult ParallelOp::verifyRegions() {
3091 auto distChildOps = getOps<DistributeOp>();
3092 int numDistChildOps = std::distance(distChildOps.begin(), distChildOps.end());
3093 if (numDistChildOps > 1)
3095 <<
"multiple 'omp.distribute' nested inside of 'omp.parallel'";
3097 if (numDistChildOps == 1) {
3100 <<
"'omp.composite' attribute missing from composite operation";
3102 auto *ompDialect =
getContext()->getLoadedDialect<OpenMPDialect>();
3103 Operation &distributeOp = **distChildOps.begin();
3105 if (&childOp == &distributeOp || ompDialect != childOp.getDialect())
3109 return emitError() <<
"unexpected OpenMP operation inside of composite "
3111 << childOp.getName();
3113 }
else if (isComposite()) {
3115 <<
"'omp.composite' attribute present in non-composite operation";
3132 const TeamsOperands &clauses) {
3136 builder, state, clauses.allocateVars, clauses.allocatorVars,
3139 clauses.dynGroupprivateAccessGroup, clauses.dynGroupprivateFallback,
3140 clauses.dynGroupprivateSize, clauses.ifExpr, clauses.numTeamsLower,
3141 clauses.numTeamsUpperVars, {},
nullptr,
3142 false, clauses.reductionMod,
3143 clauses.reductionVars,
3145 makeArrayAttr(ctx, clauses.reductionSyms), clauses.threadLimitVars);
3152 if (numTeamsLower) {
3153 if (numTeamsUpperVars.size() != 1)
3155 "expected exactly one num_teams upper bound when lower bound is "
3159 "expected num_teams upper bound and lower bound to be "
3166LogicalResult TeamsOp::verify() {
3173 auto parentTarget = llvm::dyn_cast_if_present<TargetOp>(op->
getParentOp());
3175 return emitError(
"expected to be nested inside of omp.target or not nested "
3176 "in any OpenMP dialect operations");
3180 this->getNumTeamsUpperVars())))
3184 parentTarget.getKernelType() == TargetExecMode::spmd_no_loop &&
3185 (getNumTeamsLower() || !getNumTeamsUpperVars().empty()))
3186 return emitOpError() <<
"'num_teams' not allowed in SPMD-no-loop kernels";
3189 getOperation(), getAllocateVars(), getAllocatorVars(),
3190 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3191 getPrivateVars(), getPrivateSymsAttr())))
3195 op, getDynGroupprivateAccessGroupAttr(),
3196 getDynGroupprivateFallbackAttr(), getDynGroupprivateSize())))
3203 getReductionByref());
3211 return getParentOp().getPrivateVars();
3215 return getParentOp().getReductionVars();
3223 const SectionsOperands &clauses) {
3226 SectionsOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3231 clauses.reductionMod, clauses.reductionVars,
3236LogicalResult SectionsOp::verify() {
3238 return emitOpError() <<
"cannot be a non-innermost combined construct leaf";
3241 getOperation(), getAllocateVars(), getAllocatorVars(),
3242 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3243 getPrivateVars(), getPrivateSymsAttr())))
3247 getReductionByref());
3250LogicalResult SectionsOp::verifyRegions() {
3251 for (
auto &inst : *getRegion().begin()) {
3252 if (!(isa<SectionOp>(inst) || isa<TerminatorOp>(inst))) {
3253 return emitOpError()
3254 <<
"expected omp.section op or terminator op inside region";
3266 const ScopeOperands &clauses) {
3268 ScopeOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3271 clauses.nowait, clauses.privateVars,
3273 clauses.privateNeedsBarrier, clauses.reductionMod,
3274 clauses.reductionVars,
3279LogicalResult ScopeOp::verify() {
3281 getOperation(), getAllocateVars(), getAllocatorVars(),
3282 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3283 getPrivateVars(), getPrivateSymsAttr(),
3291 getReductionByref());
3299 const SingleOperands &clauses) {
3302 SingleOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3305 clauses.copyprivateVars,
3306 makeArrayAttr(ctx, clauses.copyprivateSyms), clauses.nowait,
3311LogicalResult SingleOp::verify() {
3313 getOperation(), getAllocateVars(), getAllocatorVars(),
3314 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3315 getPrivateVars(), getPrivateSymsAttr())))
3319 getCopyprivateSyms());
3327 const WorkshareOperands &clauses) {
3328 WorkshareOp::build(builder, state, clauses.nowait);
3331LogicalResult WorkshareOp::verify() {
3333 return emitOpError() <<
"cannot be a non-innermost combined construct leaf";
3342LogicalResult WorkshareLoopWrapperOp::verifyRegions() {
3343 if (isa_and_nonnull<LoopWrapperInterface>((*this)->getParentOp()) ||
3345 return emitOpError() <<
"expected to be a standalone loop wrapper";
3354LogicalResult LoopWrapperInterface::verifyImpl() {
3358 return emitOpError() <<
"loop wrapper must also have the `NoTerminator` "
3359 "and `SingleBlock` traits";
3362 return emitOpError() <<
"loop wrapper does not contain exactly one region";
3365 if (range_size(region.
getOps()) != 1)
3366 return emitOpError()
3367 <<
"loop wrapper does not contain exactly one nested op";
3370 if (!isa<LoopNestOp, LoopWrapperInterface>(firstOp))
3371 return emitOpError() <<
"nested in loop wrapper is not another loop "
3372 "wrapper or `omp.loop_nest`";
3381Operation *ComposableOpInterface::findCapturedOp() {
3385 if (
auto wrapperOp = dyn_cast<LoopWrapperInterface>(op))
3386 return wrapperOp.getWrappedLoop();
3391 if (!isCombined() && !isComposite())
3396 if (
auto wrapperOp = dyn_cast<LoopWrapperInterface>(&nestedOp))
3397 return wrapperOp.getWrappedLoop();
3399 if (
auto composableOp = dyn_cast<ComposableOpInterface>(&nestedOp))
3400 return composableOp.findCapturedOp();
3409LogicalResult ComposableOpInterface::verifyImpl() {
3413 return emitOpError() <<
"composable ops must have a single region";
3415 if (isComposite() && !isa<LoopWrapperInterface, ParallelOp>(op))
3416 return emitOpError() <<
"non-loop wrapper cannot be composite";
3422 auto count = llvm::count_if(
3424 if (isa<ComposableOpInterface, LoopWrapperInterface>(op)) {
3445 return emitOpError()
3446 <<
"multiple eligible child ops found in combined op";
3457 if (successor->isReachable(parentBlock))
3458 return emitOpError() <<
"nested combined child op is part of a loop";
3462 !domInfo.
dominates(parentBlock, &block))
3463 return emitOpError()
3464 <<
"nested combined child op doesn't unconditionally execute";
3474 const LoopOperands &clauses) {
3477 LoopOp::build(builder, state, clauses.bindKind, clauses.privateVars,
3479 clauses.privateNeedsBarrier, clauses.order, clauses.orderMod,
3480 clauses.reductionMod, clauses.reductionVars,
3485LogicalResult LoopOp::verify() {
3490 getReductionByref());
3493LogicalResult LoopOp::verifyRegions() {
3494 if (llvm::isa_and_nonnull<LoopWrapperInterface>((*this)->getParentOp()) ||
3496 return emitOpError() <<
"expected to be a standalone loop wrapper";
3507 build(builder, state, {}, {},
3512 false,
nullptr,
nullptr,
3513 nullptr, {},
nullptr,
3524 const WsloopOperands &clauses) {
3527 builder, state, clauses.allocateVars, clauses.allocatorVars,
3530 clauses.linearVars, clauses.linearStepVars, clauses.linearVarTypes,
3531 clauses.linearModifiers, clauses.nowait, clauses.order, clauses.orderMod,
3532 clauses.ordered, clauses.privateVars,
3533 makeArrayAttr(ctx, clauses.privateSyms), clauses.privateNeedsBarrier,
3534 clauses.reductionMod, clauses.reductionVars,
3536 makeArrayAttr(ctx, clauses.reductionSyms), clauses.scheduleKind,
3537 clauses.scheduleChunk, clauses.scheduleMod, clauses.scheduleSimd);
3540LogicalResult WsloopOp::verify() {
3542 getOperation(), getAllocateVars(), getAllocatorVars(),
3543 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3544 getPrivateVars(), getPrivateSymsAttr())))
3550 if (getLinearVars().size() &&
3551 getLinearVarTypes().value().size() != getLinearVars().size())
3552 return emitError() <<
"Ill-formed type attributes for linear variables";
3558 getReductionByref());
3561LogicalResult WsloopOp::verifyRegions() {
3562 bool isCompositeChildLeaf =
3563 llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp());
3565 if (LoopWrapperInterface nested = getNestedWrapper()) {
3568 <<
"'omp.composite' attribute missing from composite wrapper";
3572 if (!isa<SimdOp>(nested))
3573 return emitError() <<
"only supported nested wrapper is 'omp.simd'";
3575 }
else if (isComposite() && !isCompositeChildLeaf) {
3577 <<
"'omp.composite' attribute present in non-composite wrapper";
3578 }
else if (!isComposite() && isCompositeChildLeaf) {
3580 <<
"'omp.composite' attribute missing from composite wrapper";
3591 const SimdOperands &clauses) {
3593 SimdOp::build(builder, state, clauses.alignedVars,
3595 clauses.linearVars, clauses.linearStepVars,
3596 clauses.linearVarTypes, clauses.linearModifiers,
3597 clauses.nontemporalVars, clauses.order, clauses.orderMod,
3598 clauses.privateVars,
makeArrayAttr(ctx, clauses.privateSyms),
3599 clauses.privateNeedsBarrier, clauses.reductionMod,
3600 clauses.reductionVars,
3606LogicalResult SimdOp::verify() {
3607 if (getSimdlen().has_value() && getSafelen().has_value() &&
3608 getSimdlen().value() > getSafelen().value())
3609 return emitOpError()
3610 <<
"simdlen clause and safelen clause are both present, but the "
3611 "simdlen value is not less than or equal to safelen value";
3623 bool isCompositeChildLeaf =
3624 llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp());
3626 if (!isComposite() && isCompositeChildLeaf)
3628 <<
"'omp.composite' attribute missing from composite wrapper";
3630 if (isComposite() && !isCompositeChildLeaf)
3632 <<
"'omp.composite' attribute present in non-composite wrapper";
3636 std::optional<ArrayAttr> privateSyms = getPrivateSyms();
3638 for (
const Attribute &sym : *privateSyms) {
3639 auto symRef = cast<SymbolRefAttr>(sym);
3640 omp::PrivateClauseOp privatizer =
3642 getOperation(), symRef);
3644 return emitError() <<
"Cannot find privatizer '" << symRef <<
"'";
3645 if (privatizer.getDataSharingType() ==
3646 DataSharingClauseType::FirstPrivate)
3647 return emitError() <<
"FIRSTPRIVATE cannot be used with SIMD";
3654 if (getLinearVars().size() &&
3655 getLinearVarTypes().value().size() != getLinearVars().size())
3656 return emitError() <<
"Ill-formed type attributes for linear variables";
3661 for (
Value var : getLinearVars()) {
3662 if (privateVars.contains(var) || reductionVars.contains(var))
3663 return emitOpError()
3664 <<
"linear variables cannot appear in other data-sharing clauses";
3670LogicalResult SimdOp::verifyRegions() {
3671 if (getNestedWrapper())
3672 return emitOpError() <<
"must wrap an 'omp.loop_nest' directly";
3682 const DistributeOperands &clauses) {
3683 DistributeOp::build(
3684 builder, state, clauses.allocateVars, clauses.allocatorVars,
3687 clauses.allocatePrivateIndices),
3688 clauses.distScheduleStatic, clauses.distScheduleChunkSize, clauses.order,
3689 clauses.orderMod, clauses.privateVars,
3691 clauses.privateNeedsBarrier);
3694LogicalResult DistributeOp::verify() {
3695 if (this->getDistScheduleChunkSize() && !this->getDistScheduleStatic())
3696 return emitOpError() <<
"chunk size set without "
3697 "dist_schedule_static being present";
3700 getOperation(), getAllocateVars(), getAllocatorVars(),
3701 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3702 getPrivateVars(), getPrivateSymsAttr())))
3711LogicalResult DistributeOp::verifyRegions() {
3712 if (LoopWrapperInterface nested = getNestedWrapper()) {
3715 <<
"'omp.composite' attribute missing from composite wrapper";
3718 if (isa<WsloopOp>(nested)) {
3720 if (!llvm::dyn_cast_if_present<ParallelOp>(parentOp) ||
3721 !cast<ComposableOpInterface>(parentOp).isComposite()) {
3722 return emitError() <<
"an 'omp.wsloop' nested wrapper is only allowed "
3723 "when a composite 'omp.parallel' is the direct "
3726 }
else if (!isa<SimdOp>(nested))
3727 return emitError() <<
"only supported nested wrappers are 'omp.simd' and "
3729 }
else if (isComposite()) {
3731 <<
"'omp.composite' attribute present in non-composite wrapper";
3742 const DeclareMapperInfoOperands &clauses) {
3743 DeclareMapperInfoOp::build(builder, state, clauses.mapVars,
3744 clauses.mapIterated);
3747LogicalResult DeclareMapperInfoOp::verify() {
3751LogicalResult DeclareMapperOp::verifyRegions() {
3752 if (!llvm::isa_and_present<DeclareMapperInfoOp>(
3753 getRegion().getBlocks().front().getTerminator()))
3754 return emitOpError() <<
"expected terminator to be a DeclareMapperInfoOp";
3763LogicalResult DeclareReductionOp::verifyRegions() {
3764 if (!getAllocRegion().empty()) {
3765 for (YieldOp yieldOp : getAllocRegion().getOps<YieldOp>()) {
3766 if (yieldOp.getResults().size() != 1 ||
3767 yieldOp.getResults().getTypes()[0] !=
getType())
3768 return emitOpError() <<
"expects alloc region to yield a value "
3769 "of the reduction type";
3773 if (getInitializerRegion().empty())
3774 return emitOpError() <<
"expects non-empty initializer region";
3775 Block &initializerEntryBlock = getInitializerRegion().
front();
3778 if (!getAllocRegion().empty())
3779 return emitOpError() <<
"expects two arguments to the initializer region "
3780 "when an allocation region is used";
3782 if (getAllocRegion().empty())
3783 return emitOpError() <<
"expects one argument to the initializer region "
3784 "when no allocation region is used";
3786 return emitOpError()
3787 <<
"expects one or two arguments to the initializer region";
3791 if (arg.getType() !=
getType())
3792 return emitOpError() <<
"expects initializer region argument to match "
3793 "the reduction type";
3795 for (YieldOp yieldOp : getInitializerRegion().getOps<YieldOp>()) {
3796 if (yieldOp.getResults().size() != 1 ||
3797 yieldOp.getResults().getTypes()[0] !=
getType())
3798 return emitOpError() <<
"expects initializer region to yield a value "
3799 "of the reduction type";
3802 if (getReductionRegion().empty())
3803 return emitOpError() <<
"expects non-empty reduction region";
3804 Block &reductionEntryBlock = getReductionRegion().
front();
3809 return emitOpError() <<
"expects reduction region with two arguments of "
3810 "the reduction type";
3811 for (YieldOp yieldOp : getReductionRegion().getOps<YieldOp>()) {
3812 if (yieldOp.getResults().size() != 1 ||
3813 yieldOp.getResults().getTypes()[0] !=
getType())
3814 return emitOpError() <<
"expects reduction region to yield a value "
3815 "of the reduction type";
3818 if (!getAtomicReductionRegion().empty()) {
3819 Block &atomicReductionEntryBlock = getAtomicReductionRegion().
front();
3823 return emitOpError() <<
"expects atomic reduction region with two "
3824 "arguments of the same type";
3825 auto ptrType = llvm::dyn_cast<PointerLikeType>(
3828 (ptrType.getElementType() && ptrType.getElementType() !=
getType()))
3829 return emitOpError() <<
"expects atomic reduction region arguments to "
3830 "be accumulators containing the reduction type";
3833 if (getCleanupRegion().empty())
3835 Block &cleanupEntryBlock = getCleanupRegion().
front();
3838 return emitOpError() <<
"expects cleanup region with one argument "
3839 "of the reduction type";
3849 const TaskOperands &clauses) {
3851 TaskOp::build(builder, state, clauses.iterated, clauses.affinityVars,
3852 clauses.allocateVars, clauses.allocatorVars,
3855 makeArrayAttr(ctx, clauses.dependKinds), clauses.dependVars,
3857 clauses.dependIterated, clauses.final, clauses.ifExpr,
3858 clauses.inReductionVars,
3860 makeArrayAttr(ctx, clauses.inReductionSyms), clauses.mergeable,
3861 clauses.priority, clauses.privateVars,
3863 clauses.privateNeedsBarrier, clauses.threadset, clauses.untied,
3864 clauses.eventHandle);
3867LogicalResult TaskOp::verify() {
3869 getOperation(), getAllocateVars(), getAllocatorVars(),
3870 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3871 getPrivateVars(), getPrivateSymsAttr())))
3874 LogicalResult verifyDependVars =
3876 getDependIteratedKinds(), getDependIterated());
3877 if (
failed(verifyDependVars))
3878 return verifyDependVars;
3884 getInReductionVars(), getInReductionByref());
3892 const TaskgroupOperands &clauses) {
3894 TaskgroupOp::build(builder, state, clauses.allocateVars,
3895 clauses.allocatorVars,
3898 clauses.taskReductionVars,
3903LogicalResult TaskgroupOp::verify() {
3905 getOperation(), getAllocateVars(), getAllocatorVars(),
3906 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr())))
3910 getTaskReductionVars(),
3911 getTaskReductionByref());
3919 const TaskloopContextOperands &clauses) {
3921 TaskloopContextOp::build(
3922 builder, state, clauses.allocateVars, clauses.allocatorVars,
3925 clauses.grainsizeMod, clauses.grainsize, clauses.ifExpr,
3926 clauses.inReductionVars,
3928 makeArrayAttr(ctx, clauses.inReductionSyms), clauses.mergeable,
3929 clauses.nogroup, clauses.numTasksMod, clauses.numTasks, clauses.priority,
3930 clauses.privateVars,
3932 clauses.privateNeedsBarrier, clauses.reductionMod, clauses.reductionVars,
3934 makeArrayAttr(ctx, clauses.reductionSyms), clauses.threadset,
3936 state.
addAttribute(
"omp.combined", UnitAttr::get(ctx));
3939TaskloopWrapperOp TaskloopContextOp::getLoopOp() {
3940 return cast<TaskloopWrapperOp>(
3942 return isa<TaskloopWrapperOp>(op);
3946LogicalResult TaskloopContextOp::verify() {
3950 getOperation(), getAllocateVars(), getAllocatorVars(),
3951 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3952 getPrivateVars(), getPrivateSymsAttr())))
3956 getReductionVars(), getReductionByref())) ||
3958 getInReductionVars(),
3959 getInReductionByref())))
3962 if (!getReductionVars().empty() && getNogroup())
3963 return emitError(
"if a reduction clause is present on the taskloop "
3964 "directive, the nogroup clause must not be specified");
3965 for (
auto var : getReductionVars()) {
3966 if (llvm::is_contained(getInReductionVars(), var))
3967 return emitError(
"the same list item cannot appear in both a reduction "
3968 "and an in_reduction clause");
3971 if (getGrainsize() && getNumTasks()) {
3973 "the grainsize clause and num_tasks clause are mutually exclusive and "
3974 "may not appear on the same taskloop directive");
3982 return emitOpError(
"must always contain the 'omp.combined' attribute");
3987LogicalResult TaskloopContextOp::verifyRegions() {
3988 Region ®ion = getRegion();
3990 return isa<TaskloopWrapperOp>(op);
3992 if (loopWrapperIt == region.
front().
end())
3993 return emitOpError()
3994 <<
"expected a TaskloopWrapperOp directly nested in the region";
3996 auto loopWrapperOp = cast<TaskloopWrapperOp>(*loopWrapperIt);
3997 auto loopNestOp = dyn_cast<LoopNestOp>(loopWrapperOp.getWrappedLoop());
4003 std::function<
bool(
Value)> isValidBoundValue = [&](
Value value) ->
bool {
4004 Region *valueRegion = value.getParentRegion();
4010 Operation *defOp = value.getDefiningOp();
4014 return llvm::all_of(defOp->
getOperands(), isValidBoundValue);
4016 auto hasUnsupportedTaskloopLocalBound = [&](
OperandRange range) ->
bool {
4017 return llvm::any_of(range,
4018 [&](
Value value) {
return !isValidBoundValue(value); });
4021 if (hasUnsupportedTaskloopLocalBound(loopNestOp.getLoopLowerBounds()) ||
4022 hasUnsupportedTaskloopLocalBound(loopNestOp.getLoopUpperBounds()) ||
4023 hasUnsupportedTaskloopLocalBound(loopNestOp.getLoopSteps())) {
4024 return emitOpError()
4025 <<
"expects loop bounds and steps to be defined outside of the "
4026 "taskloop.context region or by pure, regionless operations "
4027 "that do not depend on block arguments";
4038 const TaskloopWrapperOperands &clauses) {
4039 TaskloopWrapperOp::build(builder, state);
4042TaskloopContextOp TaskloopWrapperOp::getTaskloopContext() {
4043 return dyn_cast<TaskloopContextOp>(getOperation()->getParentOp());
4046LogicalResult TaskloopWrapperOp::verify() {
4047 TaskloopContextOp context = getTaskloopContext();
4049 return emitOpError() <<
"expected to be nested in a taskloop context op";
4053LogicalResult TaskloopWrapperOp::verifyRegions() {
4054 if (LoopWrapperInterface nested = getNestedWrapper()) {
4057 <<
"'omp.composite' attribute missing from composite wrapper";
4061 if (!isa<SimdOp>(nested))
4062 return emitError() <<
"only supported nested wrapper is 'omp.simd'";
4063 }
else if (isComposite()) {
4065 <<
"'omp.composite' attribute present in non-composite wrapper";
4089 for (
auto &iv : ivs)
4090 iv.type = loopVarType;
4095 result.addAttribute(
"loop_inclusive", UnitAttr::get(ctx));
4111 "collapse_num_loops",
4116 auto parseTiles = [&]() -> ParseResult {
4120 tiles.push_back(
tile);
4129 if (tiles.size() > 0)
4148 Region ®ion = getRegion();
4150 p <<
" (" << args <<
") : " << args[0].getType() <<
" = ("
4151 << getLoopLowerBounds() <<
") to (" << getLoopUpperBounds() <<
") ";
4152 if (getLoopInclusive())
4154 p <<
"step (" << getLoopSteps() <<
") ";
4155 if (
int64_t numCollapse = getCollapseNumLoops())
4156 if (numCollapse > 1)
4157 p <<
"collapse(" << numCollapse <<
") ";
4160 p <<
"tiles(" << tiles.value() <<
") ";
4166 const LoopNestOperands &clauses) {
4168 LoopNestOp::build(builder, state, clauses.collapseNumLoops,
4169 clauses.loopLowerBounds, clauses.loopUpperBounds,
4170 clauses.loopSteps, clauses.loopInclusive,
4174LogicalResult LoopNestOp::verify() {
4175 if (getLoopLowerBounds().empty())
4176 return emitOpError() <<
"must represent at least one loop";
4178 if (getLoopLowerBounds().size() != getIVs().size())
4179 return emitOpError() <<
"number of range arguments and IVs do not match";
4181 for (
auto [lb, iv] : llvm::zip_equal(getLoopLowerBounds(), getIVs())) {
4182 if (lb.getType() != iv.getType())
4183 return emitOpError()
4184 <<
"range argument type does not match corresponding IV type";
4187 uint64_t numIVs = getIVs().size();
4189 if (
const auto &numCollapse = getCollapseNumLoops())
4190 if (numCollapse > numIVs)
4191 return emitOpError()
4192 <<
"collapse value is larger than the number of loops";
4195 if (tiles.value().size() > numIVs)
4196 return emitOpError() <<
"too few canonical loops for tile dimensions";
4198 if (!llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp()))
4199 return emitOpError() <<
"expects parent op to be a loop wrapper";
4204void LoopNestOp::gatherWrappers(
4207 while (
auto wrapper =
4208 llvm::dyn_cast_if_present<LoopWrapperInterface>(parent)) {
4209 wrappers.push_back(wrapper);
4218std::tuple<NewCliOp, OpOperand *, OpOperand *>
4224 return {{},
nullptr,
nullptr};
4227 "Unexpected type of cli");
4233 auto op = cast<LoopTransformationInterface>(use.getOwner());
4235 unsigned opnum = use.getOperandNumber();
4236 if (op.isGeneratee(opnum)) {
4237 assert(!gen &&
"Each CLI may have at most one def");
4239 }
else if (op.isApplyee(opnum)) {
4240 assert(!cons &&
"Each CLI may have at most one consumer");
4243 llvm_unreachable(
"Unexpected operand for a CLI");
4247 return {create, gen, cons};
4253 case llvm::omp::ProcBindKind::OMP_PROC_BIND_close:
4254 return ClauseProcBindKind::Close;
4255 case llvm::omp::ProcBindKind::OMP_PROC_BIND_master:
4256 return ClauseProcBindKind::Master;
4257 case llvm::omp::ProcBindKind::OMP_PROC_BIND_primary:
4258 return ClauseProcBindKind::Primary;
4259 case llvm::omp::ProcBindKind::OMP_PROC_BIND_spread:
4260 return ClauseProcBindKind::Spread;
4261 case llvm::omp::ProcBindKind::OMP_PROC_BIND_default:
4262 case llvm::omp::ProcBindKind::OMP_PROC_BIND_unknown:
4265 llvm_unreachable(
"unexpected proc-bind kind");
4288 std::string cliName{
"cli"};
4292 .Case([&](CanonicalLoopOp op) {
4295 .Case([&](UnrollHeuristicOp op) -> std::string {
4296 llvm_unreachable(
"heuristic unrolling does not generate a loop");
4298 .Case([&](FuseOp op) -> std::string {
4299 unsigned opnum =
generator->getOperandNumber();
4302 if (op.getFirst().has_value() && opnum != op.getFirst().value())
4303 return "canonloop_fuse";
4307 .Case([&](TileOp op) -> std::string {
4308 auto [generateesFirst, generateesCount] =
4309 op.getGenerateesODSOperandIndexAndLength();
4310 unsigned firstGrid = generateesFirst;
4311 unsigned firstIntratile = generateesFirst + generateesCount / 2;
4312 unsigned end = generateesFirst + generateesCount;
4313 unsigned opnum =
generator->getOperandNumber();
4315 if (firstGrid <= opnum && opnum < firstIntratile) {
4316 unsigned gridnum = opnum - firstGrid + 1;
4317 return (
"grid" + Twine(gridnum)).str();
4319 if (firstIntratile <= opnum && opnum < end) {
4320 unsigned intratilenum = opnum - firstIntratile + 1;
4321 return (
"intratile" + Twine(intratilenum)).str();
4323 llvm_unreachable(
"Unexpected generatee argument");
4325 .DefaultUnreachable(
"TODO: Custom name for this operation");
4328 setNameFn(
result, cliName);
4331LogicalResult NewCliOp::verify() {
4332 Value cli = getResult();
4335 "Unexpected type of cli");
4341 auto op = cast<mlir::omp::LoopTransformationInterface>(use.getOwner());
4343 unsigned opnum = use.getOperandNumber();
4344 if (op.isGeneratee(opnum)) {
4347 emitOpError(
"CLI must have at most one generator");
4349 .
append(
"first generator here:");
4351 .
append(
"second generator here:");
4356 }
else if (op.isApplyee(opnum)) {
4359 emitOpError(
"CLI must have at most one consumer");
4361 .
append(
"first consumer here:")
4365 .
append(
"second consumer here:")
4372 llvm_unreachable(
"Unexpected operand for a CLI");
4380 .
append(
"see consumer here: ")
4403 setNameFn(&getRegion().front(),
"body_entry");
4406void CanonicalLoopOp::getAsmBlockArgumentNames(
Region ®ion,
4414 p <<
'(' << getCli() <<
')';
4415 p <<
' ' << getInductionVar() <<
" : " << getInductionVar().getType()
4416 <<
" in range(" << getTripCount() <<
") ";
4426 CanonicalLoopInfoType cliType =
4427 CanonicalLoopInfoType::get(parser.
getContext());
4452 if (parser.
parseRegion(*region, {inductionVariable}))
4457 result.operands.append(cliOperand);
4463 return mlir::success();
4466LogicalResult CanonicalLoopOp::verify() {
4469 if (!getRegion().empty()) {
4470 Region ®ion = getRegion();
4473 "Canonical loop region must have exactly one argument");
4477 "Region argument must be the same type as the trip count");
4483Value CanonicalLoopOp::getInductionVar() {
return getRegion().getArgument(0); }
4485std::pair<unsigned, unsigned>
4486CanonicalLoopOp::getApplyeesODSOperandIndexAndLength() {
4491std::pair<unsigned, unsigned>
4492CanonicalLoopOp::getGenerateesODSOperandIndexAndLength() {
4493 return getODSOperandIndexAndLength(odsIndex_cli);
4507 p <<
'(' << getApplyee() <<
')';
4514 auto cliType = CanonicalLoopInfoType::get(parser.
getContext());
4537 return mlir::success();
4540std::pair<unsigned, unsigned>
4541UnrollHeuristicOp ::getApplyeesODSOperandIndexAndLength() {
4542 return getODSOperandIndexAndLength(odsIndex_applyee);
4545std::pair<unsigned, unsigned>
4546UnrollHeuristicOp::getGenerateesODSOperandIndexAndLength() {
4560 p <<
'(' << getApplyee() <<
')';
4567 auto cliType = CanonicalLoopInfoType::get(parser.
getContext());
4590 return mlir::success();
4593std::pair<unsigned, unsigned>
4594UnrollFullOp::getApplyeesODSOperandIndexAndLength() {
4595 return getODSOperandIndexAndLength(odsIndex_applyee);
4598std::pair<unsigned, unsigned>
4599UnrollFullOp::getGenerateesODSOperandIndexAndLength() {
4603LogicalResult UnrollFullOp::verify() {
4604 auto [create, gen, cons] =
decodeCli(getApplyee());
4606 return emitOpError() <<
"applyee CLI has no generator";
4610 if (
auto loop = dyn_cast<CanonicalLoopOp>(gen->getOwner())) {
4612 return emitOpError() <<
"applyee loop must have a constant trip count";
4624 uint64_t unrollFactor) {
4631 p <<
'(' << getApplyee() <<
')';
4634 attrs.emplace_back(getUnrollFactorAttrName(), getUnrollFactorAttr());
4641 auto cliType = CanonicalLoopInfoType::get(parser.
getContext());
4658 return mlir::success();
4661std::pair<unsigned, unsigned>
4662UnrollPartialOp::getApplyeesODSOperandIndexAndLength() {
4663 return getODSOperandIndexAndLength(odsIndex_applyee);
4666std::pair<unsigned, unsigned>
4667UnrollPartialOp::getGenerateesODSOperandIndexAndLength() {
4678 if (!generatees.empty())
4679 p <<
'(' << llvm::interleaved(generatees) <<
')';
4681 if (!applyees.empty())
4682 p <<
" <- (" << llvm::interleaved(applyees) <<
')';
4724 bool isOnlyCanonLoops =
true;
4726 for (
Value applyee : op.getApplyees()) {
4727 auto [create, gen, cons] =
decodeCli(applyee);
4730 return op.emitOpError() <<
"applyee CLI has no generator";
4732 auto loop = dyn_cast_or_null<CanonicalLoopOp>(gen->getOwner());
4733 canonLoops.push_back(loop);
4735 isOnlyCanonLoops =
false;
4740 if (!isOnlyCanonLoops)
4744 for (
auto i : llvm::seq<int>(1, canonLoops.size())) {
4745 auto parentLoop = canonLoops[i - 1];
4746 auto loop = canonLoops[i];
4748 if (parentLoop.getOperation() != loop.getOperation()->getParentOp())
4749 return op.emitOpError()
4750 <<
"tiled loop nest must be nested within each other";
4752 parentIVs.insert(parentLoop.getInductionVar());
4757 bool isPerfectlyNested = [&]() {
4758 auto &parentBody = parentLoop.getRegion();
4759 if (!parentBody.hasOneBlock())
4761 auto &parentBlock = parentBody.getBlocks().
front();
4763 auto nestedLoopIt = parentBlock.
begin();
4764 if (nestedLoopIt == parentBlock.
end() ||
4765 (&*nestedLoopIt != loop.getOperation()))
4768 auto termIt = std::next(nestedLoopIt);
4769 if (termIt == parentBlock.
end() || !isa<TerminatorOp>(termIt))
4772 if (std::next(termIt) != parentBlock.
end())
4777 if (!isPerfectlyNested)
4778 return op.emitOpError() <<
"tiled loop nest must be perfectly nested";
4780 if (parentIVs.contains(loop.getTripCount()))
4781 return op.emitOpError() <<
"tiled loop nest must be rectangular";
4798LogicalResult TileOp::verify() {
4799 if (getApplyees().empty())
4800 return emitOpError() <<
"must apply to at least one loop";
4802 if (getSizes().size() != getApplyees().size())
4803 return emitOpError() <<
"there must be one tile size for each applyee";
4805 if (!getGeneratees().empty() &&
4806 2 * getSizes().size() != getGeneratees().size())
4807 return emitOpError()
4808 <<
"expecting two times the number of generatees than applyees";
4813std::pair<unsigned, unsigned> TileOp ::getApplyeesODSOperandIndexAndLength() {
4814 return getODSOperandIndexAndLength(odsIndex_applyees);
4817std::pair<unsigned, unsigned> TileOp::getGenerateesODSOperandIndexAndLength() {
4818 return getODSOperandIndexAndLength(odsIndex_generatees);
4828 if (!generatees.empty())
4829 p <<
'(' << llvm::interleaved(generatees) <<
')';
4831 if (!applyees.empty())
4832 p <<
" <- (" << llvm::interleaved(applyees) <<
')';
4835LogicalResult FuseOp::verify() {
4836 if (getApplyees().size() < 2)
4837 return emitOpError() <<
"must apply to at least two loops";
4839 if (getFirst().has_value() && getCount().has_value()) {
4840 int64_t first = getFirst().value();
4841 int64_t count = getCount().value();
4842 if ((
unsigned)(first + count - 1) > getApplyees().size())
4843 return emitOpError() <<
"the numbers of applyees must be at least first "
4844 "minus one plus count attributes";
4845 if (!getGeneratees().empty() &&
4846 getGeneratees().size() != getApplyees().size() + 1 - count)
4847 return emitOpError() <<
"the number of generatees must be the number of "
4848 "aplyees plus one minus count";
4851 if (!getGeneratees().empty() && getGeneratees().size() != 1)
4852 return emitOpError()
4853 <<
"in a complete fuse the number of generatees must be exactly 1";
4855 for (
auto &&applyee : getApplyees()) {
4856 auto [create, gen, cons] =
decodeCli(applyee);
4859 return emitOpError() <<
"applyee CLI has no generator";
4860 auto loop = dyn_cast_or_null<CanonicalLoopOp>(gen->getOwner());
4862 return emitOpError()
4863 <<
"currently only supports omp.canonical_loop as applyee";
4867std::pair<unsigned, unsigned> FuseOp::getApplyeesODSOperandIndexAndLength() {
4868 return getODSOperandIndexAndLength(odsIndex_applyees);
4871std::pair<unsigned, unsigned> FuseOp::getGenerateesODSOperandIndexAndLength() {
4872 return getODSOperandIndexAndLength(odsIndex_generatees);
4880 const CriticalDeclareOperands &clauses) {
4881 CriticalDeclareOp::build(builder, state, clauses.symName,
4882 clauses.symVisibility, clauses.hint);
4885LogicalResult CriticalDeclareOp::verify() {
4889LogicalResult CriticalOp::verify() {
4890 SymbolRefAttr currentName = getNameAttr();
4892 CriticalOp parentCritical = (*this)->getParentOfType<CriticalOp>();
4894 while (parentCritical) {
4895 SymbolRefAttr parentName = parentCritical.getNameAttr();
4897 if (currentName == parentName) {
4899 return emitOpError() <<
"cannot be nested inside another omp.critical "
4900 "region with the same name ("
4901 << currentName <<
")";
4903 return emitOpError() <<
"cannot be nested inside another unnamed "
4904 "omp.critical region";
4908 parentCritical = parentCritical->getParentOfType<CriticalOp>();
4915 if (getNameAttr()) {
4916 SymbolRefAttr symbolRef = getNameAttr();
4920 return emitOpError() <<
"expected symbol reference " << symbolRef
4921 <<
" to point to a critical declaration";
4932LogicalResult ErrorOp::verify() {
4933 if (getMessage() && getMessageExpr())
4934 return emitOpError() <<
"the message must be provided either as a constant "
4935 "`message` attribute or as a `message_expr` "
4936 "operand, but not both";
4953 return op.
emitOpError() <<
"must be nested inside of a loop";
4957 if (
auto wsloopOp = dyn_cast<WsloopOp>(wrapper)) {
4958 IntegerAttr orderedAttr = wsloopOp.getOrderedAttr();
4960 return op.
emitOpError() <<
"the enclosing worksharing-loop region must "
4961 "have an ordered clause";
4963 if (hasRegion && orderedAttr.getInt() != 0)
4964 return op.
emitOpError() <<
"the enclosing loop's ordered clause must not "
4965 "have a parameter present";
4967 if (!hasRegion && orderedAttr.getInt() == 0)
4968 return op.
emitOpError() <<
"the enclosing loop's ordered clause must "
4969 "have a parameter present";
4970 }
else if (!isa<SimdOp>(wrapper)) {
4971 return op.
emitOpError() <<
"must be nested inside of a worksharing, simd "
4972 "or worksharing simd loop";
4978 const OrderedOperands &clauses) {
4979 OrderedOp::build(builder, state, clauses.doacrossDependType,
4980 clauses.doacrossNumLoops, clauses.doacrossDependVars);
4983LogicalResult OrderedOp::verify() {
4987 auto wrapper = (*this)->getParentOfType<WsloopOp>();
4988 if (!wrapper || *wrapper.getOrdered() != *getDoacrossNumLoops())
4989 return emitOpError() <<
"number of variables in depend clause does not "
4990 <<
"match number of iteration variables in the "
4997 const OrderedRegionOperands &clauses) {
4998 OrderedRegionOp::build(builder, state, clauses.parLevelSimd);
5008 const TaskwaitOperands &clauses) {
5024LogicalResult AtomicReadOp::verify() {
5025 if (verifyCommon().
failed())
5026 return mlir::failure();
5029 if (
auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5030 if (
Attribute verAttr = moduleOp->getDiscardableAttr(
"omp.version"))
5031 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5033 if (
auto mo = getMemoryOrder()) {
5034 if (*mo == ClauseMemoryOrderKind::Release) {
5035 return emitError(
"memory-order must not be release for atomic reads");
5037 if (*mo == ClauseMemoryOrderKind::Acq_rel) {
5040 return emitError(
"memory-order must not be acq_rel for atomic reads");
5050LogicalResult AtomicWriteOp::verify() {
5051 if (verifyCommon().
failed())
5052 return mlir::failure();
5055 if (
auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5056 if (
Attribute verAttr = moduleOp->getDiscardableAttr(
"omp.version"))
5057 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5059 if (
auto mo = getMemoryOrder()) {
5060 if (*mo == ClauseMemoryOrderKind::Acquire) {
5061 return emitError(
"memory-order must not be acquire for atomic writes");
5063 if (*mo == ClauseMemoryOrderKind::Acq_rel) {
5066 return emitError(
"memory-order must not be acq_rel for atomic writes");
5076LogicalResult AtomicUpdateOp::canonicalize(AtomicUpdateOp op,
5082 if (
Value writeVal = op.getWriteOpVal()) {
5084 op, op.getX(), writeVal, op.getHintAttr(), op.getMemoryOrderAttr());
5090LogicalResult AtomicUpdateOp::verify() {
5091 if (verifyCommon().
failed())
5092 return mlir::failure();
5095 if (
auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5096 if (
Attribute verAttr = moduleOp->getDiscardableAttr(
"omp.version"))
5097 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5099 if (
auto mo = getMemoryOrder()) {
5100 if (*mo == ClauseMemoryOrderKind::Acq_rel ||
5101 *mo == ClauseMemoryOrderKind::Acquire) {
5105 "memory-order must not be acq_rel or acquire for atomic updates");
5112LogicalResult AtomicUpdateOp::verifyRegions() {
return verifyRegionsCommon(); }
5118AtomicReadOp AtomicCaptureOp::getAtomicReadOp() {
5119 if (
auto op = dyn_cast<AtomicReadOp>(getFirstOp()))
5121 return dyn_cast<AtomicReadOp>(getSecondOp());
5124AtomicWriteOp AtomicCaptureOp::getAtomicWriteOp() {
5125 if (
auto op = dyn_cast<AtomicWriteOp>(getFirstOp()))
5127 return dyn_cast<AtomicWriteOp>(getSecondOp());
5130AtomicUpdateOp AtomicCaptureOp::getAtomicUpdateOp() {
5131 if (
auto op = dyn_cast<AtomicUpdateOp>(getFirstOp()))
5133 return dyn_cast<AtomicUpdateOp>(getSecondOp());
5136AtomicCompareOp AtomicCaptureOp::getAtomicCompareOp() {
5137 if (
auto op = dyn_cast<AtomicCompareOp>(getFirstOp()))
5139 return dyn_cast<AtomicCompareOp>(getSecondOp());
5142LogicalResult AtomicCaptureOp::verify() {
5146LogicalResult AtomicCaptureOp::verifyRegions() {
5147 if (verifyRegionsCommon().
failed())
5148 return mlir::failure();
5150 if (getFirstOp()->getInherentAttr(
"hint").value_or(
Attribute{}) ||
5151 getSecondOp()->getInherentAttr(
"hint").value_or(
Attribute{}))
5153 "operations inside capture region must not have hint clause");
5155 if (getFirstOp()->getInherentAttr(
"memory_order").value_or(
Attribute{}) ||
5156 getSecondOp()->getInherentAttr(
"memory_order").value_or(
Attribute{}))
5158 "operations inside capture region must not have memory_order clause");
5166LogicalResult AtomicCompareOp::verify() {
5167 if (verifyCommon().
failed())
5168 return mlir::failure();
5172 if (
auto failOrder = getFailMemoryOrder()) {
5173 if (*failOrder != ClauseMemoryOrderKind::Seq_cst &&
5174 *failOrder != ClauseMemoryOrderKind::Acquire &&
5175 *failOrder != ClauseMemoryOrderKind::Relaxed)
5177 "fail_memory_order must be 'seq_cst', 'acquire' or 'relaxed'");
5182LogicalResult AtomicCompareOp::verifyRegions() {
5183 if (verifyRegionsCommon().
failed())
5184 return mlir::failure();
5186 if (verifyOperator().
failed())
5187 return mlir::failure();
5192 if (!terminator || !isa<YieldOp>(terminator))
5193 return emitOpError(
"region must be terminated with omp.yield");
5203 const CancelOperands &clauses) {
5204 CancelOp::build(builder, state, clauses.cancelDirective, clauses.ifExpr);
5217LogicalResult CancelOp::verify() {
5218 ClauseCancellationConstructType cct = getCancelDirective();
5221 if (!structuralParent)
5222 return emitOpError() <<
"Orphaned cancel construct";
5224 if ((cct == ClauseCancellationConstructType::Parallel) &&
5225 !mlir::isa<ParallelOp>(structuralParent)) {
5226 return emitOpError() <<
"cancel parallel must appear "
5227 <<
"inside a parallel region";
5229 if (cct == ClauseCancellationConstructType::Loop) {
5232 auto wsloopOp = mlir::dyn_cast<WsloopOp>(structuralParent->
getParentOp());
5235 return emitOpError()
5236 <<
"cancel loop must appear inside a worksharing-loop region";
5238 if (wsloopOp.getNowaitAttr()) {
5239 return emitError() <<
"A worksharing construct that is canceled "
5240 <<
"must not have a nowait clause";
5242 if (wsloopOp.getOrderedAttr()) {
5243 return emitError() <<
"A worksharing construct that is canceled "
5244 <<
"must not have an ordered clause";
5247 }
else if (cct == ClauseCancellationConstructType::Sections) {
5251 mlir::dyn_cast<SectionsOp>(structuralParent->
getParentOp());
5253 return emitOpError() <<
"cancel sections must appear "
5254 <<
"inside a sections region";
5256 if (sectionsOp.getNowait()) {
5257 return emitError() <<
"A sections construct that is canceled "
5258 <<
"must not have a nowait clause";
5261 if ((cct == ClauseCancellationConstructType::Taskgroup) &&
5262 (!mlir::isa<omp::TaskOp>(structuralParent) &&
5263 !mlir::isa<omp::TaskloopWrapperOp>(structuralParent->
getParentOp()))) {
5264 return emitOpError() <<
"cancel taskgroup must appear "
5265 <<
"inside a task region";
5275 const CancellationPointOperands &clauses) {
5276 CancellationPointOp::build(builder, state, clauses.cancelDirective);
5279LogicalResult CancellationPointOp::verify() {
5280 ClauseCancellationConstructType cct = getCancelDirective();
5283 if (!structuralParent)
5284 return emitOpError() <<
"Orphaned cancellation point";
5286 if ((cct == ClauseCancellationConstructType::Parallel) &&
5287 !mlir::isa<ParallelOp>(structuralParent)) {
5288 return emitOpError() <<
"cancellation point parallel must appear "
5289 <<
"inside a parallel region";
5293 if ((cct == ClauseCancellationConstructType::Loop) &&
5294 !mlir::isa<WsloopOp>(structuralParent->
getParentOp())) {
5295 return emitOpError() <<
"cancellation point loop must appear "
5296 <<
"inside a worksharing-loop region";
5298 if ((cct == ClauseCancellationConstructType::Sections) &&
5299 !mlir::isa<omp::SectionOp>(structuralParent)) {
5300 return emitOpError() <<
"cancellation point sections must appear "
5301 <<
"inside a sections region";
5303 if ((cct == ClauseCancellationConstructType::Taskgroup) &&
5304 (!mlir::isa<omp::TaskOp>(structuralParent) &&
5305 !mlir::isa<omp::TaskloopWrapperOp>(structuralParent->
getParentOp()))) {
5306 return emitOpError() <<
"cancellation point taskgroup must appear "
5307 <<
"inside a task region";
5316LogicalResult MapBoundsOp::verify() {
5317 auto extent = getExtent();
5319 if (!extent && !upperbound)
5320 return emitError(
"expected extent or upperbound.");
5327 PrivateClauseOp::build(
5328 odsBuilder, odsState, symName,
nullptr, type,
5329 DataSharingClauseTypeAttr::get(odsBuilder.
getContext(),
5330 DataSharingClauseType::Private));
5333LogicalResult PrivateClauseOp::verifyRegions() {
5334 Type argType = getArgType();
5335 auto verifyTerminator = [&](
Operation *terminator,
5336 bool yieldsValue) -> LogicalResult {
5340 if (!llvm::isa<YieldOp>(terminator))
5342 <<
"expected exit block terminator to be an `omp.yield` op.";
5344 YieldOp yieldOp = llvm::cast<YieldOp>(terminator);
5345 TypeRange yieldedTypes = yieldOp.getResults().getTypes();
5348 if (yieldedTypes.empty())
5352 <<
"Did not expect any values to be yielded.";
5355 if (yieldedTypes.size() == 1 && yieldedTypes.front() == argType)
5359 <<
"Invalid yielded value. Expected type: " << argType
5362 if (yieldedTypes.empty())
5365 error << yieldedTypes;
5371 StringRef regionName,
5372 bool yieldsValue) -> LogicalResult {
5373 assert(!region.
empty());
5377 <<
"`" << regionName <<
"`: " <<
"expected " << expectedNumArgs
5380 for (
Block &block : region) {
5393 for (
Region *region : getRegions())
5394 for (
Type ty : region->getArgumentTypes())
5396 return emitError() <<
"Region argument type mismatch: got " << ty
5397 <<
" expected " << argType <<
".";
5400 if (!initRegion.
empty() &&
5405 DataSharingClauseType dsType = getDataSharingType();
5407 if (dsType == DataSharingClauseType::Private && !getCopyRegion().empty())
5408 return emitError(
"`private` clauses do not require a `copy` region.");
5410 if (dsType == DataSharingClauseType::FirstPrivate && getCopyRegion().empty())
5412 "`firstprivate` clauses require at least a `copy` region.");
5414 if (dsType == DataSharingClauseType::FirstPrivate &&
5419 if (!getDeallocRegion().empty() &&
5432 const MaskedOperands &clauses) {
5433 MaskedOp::build(builder, state, clauses.filteredThreadId);
5441 const DispatchOperands &clauses) {
5442 DispatchOp::build(builder, state, clauses.nocontext, clauses.novariants,
5451 const ScanOperands &clauses) {
5452 ScanOp::build(builder, state, clauses.inclusiveVars, clauses.exclusiveVars);
5455LogicalResult ScanOp::verify() {
5456 if (hasExclusiveVars() == hasInclusiveVars())
5458 "Exactly one of EXCLUSIVE or INCLUSIVE clause is expected");
5459 if (WsloopOp parentWsLoopOp = (*this)->getParentOfType<WsloopOp>()) {
5460 if (parentWsLoopOp.getReductionModAttr() &&
5461 parentWsLoopOp.getReductionModAttr().getValue() ==
5462 ReductionModifier::inscan)
5465 if (SimdOp parentSimdOp = (*this)->getParentOfType<SimdOp>()) {
5466 if (parentSimdOp.getReductionModAttr() &&
5467 parentSimdOp.getReductionModAttr().getValue() ==
5468 ReductionModifier::inscan)
5471 return emitError(
"SCAN directive needs to be enclosed within a parent "
5472 "worksharing loop construct or SIMD construct with INSCAN "
5473 "reduction modifier");
5478 std::optional<uint64_t> alignment) {
5479 if (alignment.has_value()) {
5480 if ((alignment.value() != 0) && !llvm::has_single_bit(alignment.value()))
5482 <<
"ALIGN value : " << alignment.value() <<
" must be power of 2";
5487LogicalResult AllocateDirOp::verify() {
5495LogicalResult AllocSharedMemOp::verify() {
5503LogicalResult FreeSharedMemOp::verify() {
5511LogicalResult WorkdistributeOp::verify() {
5513 return emitOpError() <<
"cannot be a non-innermost combined construct leaf";
5516 Region ®ion = getRegion();
5518 return emitOpError(
"region cannot be empty");
5521 if (entryBlock.
empty())
5522 return emitOpError(
"region must contain a structured block");
5524 bool hasTerminator =
false;
5525 for (
Block &block : region) {
5526 if (isa<TerminatorOp>(block.
back())) {
5527 if (hasTerminator) {
5528 return emitOpError(
"region must have exactly one terminator");
5530 hasTerminator =
true;
5533 if (!hasTerminator) {
5534 return emitOpError(
"region must be terminated with omp.terminator");
5538 if (isa<BarrierOp>(op)) {
5540 "explicit barriers are not allowed in workdistribute region");
5543 if (isa<ParallelOp>(op)) {
5545 "nested parallel constructs not allowed in workdistribute");
5547 if (isa<TeamsOp>(op)) {
5549 "nested teams constructs not allowed in workdistribute");
5553 if (walkResult.wasInterrupted())
5557 if (!llvm::dyn_cast<TeamsOp>(parentOp))
5558 return emitOpError(
"workdistribute must be nested under teams");
5566LogicalResult DeclareSimdOp::verify() {
5569 dyn_cast_if_present<mlir::FunctionOpInterface>((*this)->getParentOp());
5571 return emitOpError() <<
"must be nested inside a function";
5573 if (getInbranch() && getNotinbranch())
5574 return emitOpError(
"cannot have both 'inbranch' and 'notinbranch'");
5584 const DeclareSimdOperands &clauses) {
5586 DeclareSimdOp::build(odsBuilder, odsState, clauses.alignedVars,
5588 clauses.linearVars, clauses.linearStepVars,
5589 clauses.linearVarTypes, clauses.linearModifiers,
5590 clauses.notinbranch, clauses.simdlen,
5591 clauses.uniformVars);
5608 return mlir::failure();
5609 return mlir::success();
5616 for (
unsigned i = 0; i < uniformVars.size(); ++i) {
5619 p << uniformVars[i] <<
" : " << uniformTypes[i];
5634 parser, iterated, iteratedTypes, affinityVars, affinityVarTypes,
5635 [&]() -> ParseResult {
return success(); })))
5669 OpAsmParser::Argument &arg = ivArgs.emplace_back();
5670 if (parser.parseArgument(arg))
5674 if (succeeded(parser.parseOptionalColon())) {
5675 if (parser.parseType(arg.type))
5678 arg.type = parser.getBuilder().getIndexType();
5690 OpAsmParser::UnresolvedOperand lb, ub, st;
5691 if (parser.parseOperand(lb) || parser.parseKeyword(
"to") ||
5692 parser.parseOperand(ub) || parser.parseKeyword(
"step") ||
5693 parser.parseOperand(st))
5698 steps.push_back(st);
5706 if (ivArgs.size() != lbs.size())
5708 <<
"mismatch: " << ivArgs.size() <<
" variables but " << lbs.size()
5711 for (
auto &arg : ivArgs) {
5712 lbTypes.push_back(arg.type);
5713 ubTypes.push_back(arg.type);
5714 stepTypes.push_back(arg.type);
5734 for (
unsigned i = 0, e = lbs.size(); i < e; ++i) {
5737 p << lbs[i] <<
" to " << ubs[i] <<
" step " << steps[i];
5745LogicalResult IteratorOp::verify() {
5746 auto iteratedTy = llvm::dyn_cast<omp::IteratedType>(getIterated().
getType());
5748 return emitOpError() <<
"result must be omp.iterated<entry_ty>";
5750 for (
auto [lb,
ub, step] : llvm::zip_equal(
5751 getLoopLowerBounds(), getLoopUpperBounds(), getLoopSteps())) {
5753 return emitOpError() <<
"loop step must not be zero";
5757 IntegerAttr stepAttr;
5763 const APInt &lbVal = lbAttr.getValue();
5764 const APInt &ubVal = ubAttr.getValue();
5765 const APInt &stepVal = stepAttr.getValue();
5766 if (stepVal.isStrictlyPositive() && lbVal.sgt(ubVal))
5767 return emitOpError() <<
"positive loop step requires lower bound to be "
5768 "less than or equal to upper bound";
5769 if (stepVal.isNegative() && lbVal.slt(ubVal))
5770 return emitOpError() <<
"negative loop step requires lower bound to be "
5771 "greater than or equal to upper bound";
5774 Block &
b = getRegion().front();
5775 auto yield = llvm::dyn_cast<omp::YieldOp>(
b.getTerminator());
5778 return emitOpError() <<
"region must be terminated by omp.yield";
5780 if (yield.getNumOperands() != 1)
5781 return emitOpError()
5782 <<
"omp.yield in omp.iterator region must yield exactly one value";
5784 mlir::Type yieldedTy = yield.getOperand(0).getType();
5785 mlir::Type elemTy = iteratedTy.getElementType();
5787 if (yieldedTy != elemTy)
5788 return emitOpError() <<
"omp.iterated element type (" << elemTy
5789 <<
") does not match omp.yield operand type ("
5790 << yieldedTy <<
")";
5803 return emitOpError() <<
"expected symbol reference '" << getSymName()
5804 <<
"' to point to a global variable";
5806 if (isa<FunctionOpInterface>(symbol))
5807 return emitOpError() <<
"expected symbol reference '" << getSymName()
5808 <<
"' to point to a global variable, not a function";
5813#define GET_ATTRDEF_CLASSES
5814#include "mlir/Dialect/OpenMP/OpenMPOpsAttributes.cpp.inc"
5816#define GET_OP_CLASSES
5817#include "mlir/Dialect/OpenMP/OpenMPOps.cpp.inc"
5819#define GET_TYPEDEF_CLASSES
5820#include "mlir/Dialect/OpenMP/OpenMPOpsTypes.cpp.inc"
static std::optional< int64_t > getUpperBound(Value iv)
Gets the constant upper bound on an affine.for iv.
static LogicalResult verifyRegion(emitc::SwitchOp op, Region ®ion, const Twine &name)
static const mlir::GenInfo * generator
static LogicalResult verifyNontemporalClause(Operation *op, OperandRange nontemporalVars)
static DenseI64ArrayAttr makeDenseI64ArrayAttr(MLIRContext *ctx, const ArrayRef< int64_t > intArray)
static void printDependVarList(OpAsmPrinter &p, Operation *op, OperandRange dependVars, TypeRange dependTypes, std::optional< ArrayAttr > dependKinds, OperandRange iteratedVars, TypeRange iteratedTypes, std::optional< ArrayAttr > iteratedKinds)
Print Depend clause.
static ParseResult parseTargetOpRegion(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &hasDeviceAddrVars, SmallVectorImpl< Type > &hasDeviceAddrTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &hostEvalVars, SmallVectorImpl< Type > &hostEvalTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &mapVars, SmallVectorImpl< Type > &mapTypes, llvm::SmallVectorImpl< OpAsmParser::UnresolvedOperand > &privateVars, llvm::SmallVectorImpl< Type > &privateTypes, ArrayAttr &privateSyms, UnitAttr &privateNeedsBarrier, DenseI64ArrayAttr &privateMaps)
static constexpr StringRef getPrivateNeedsBarrierSpelling()
static void printHeapAllocClause(OpAsmPrinter &p, Operation *op, TypeAttr inType, ValueRange typeparams, TypeRange typeparamsTypes, ValueRange shape, TypeRange shapeTypes)
static LogicalResult verifyReductionVarList(Operation *op, std::optional< ArrayAttr > reductionSyms, OperandRange reductionVars, std::optional< ArrayRef< bool > > reductionByref)
Verifies Reduction Clause.
static ParseResult parseLinearClause(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &linearVars, SmallVectorImpl< Type > &linearTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &linearStepVars, SmallVectorImpl< Type > &linearStepTypes, ArrayAttr &linearModifiers)
linear ::= linear ( linear-list ) linear-list := linear-val | linear-val linear-list linear-val := ss...
static ParseResult parseInReductionPrivateRegion(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &inReductionVars, SmallVectorImpl< Type > &inReductionTypes, DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms, llvm::SmallVectorImpl< OpAsmParser::UnresolvedOperand > &privateVars, llvm::SmallVectorImpl< Type > &privateTypes, ArrayAttr &privateSyms, UnitAttr &privateNeedsBarrier)
static ArrayAttr makeArrayAttr(MLIRContext *context, llvm::ArrayRef< Attribute > attrs)
static ParseResult parseClauseAttr(AsmParser &parser, ClauseAttr &attr)
static void printDynGroupprivateClause(OpAsmPrinter &printer, Operation *op, AccessGroupModifierAttr modifierFirst, FallbackModifierAttr modifierSecond, Value dynGroupprivateSize, Type sizeType)
static void printAllocateAndAllocator(OpAsmPrinter &p, Operation *op, OperandRange allocateVars, TypeRange allocateTypes, OperandRange allocatorVars, TypeRange allocatorTypes)
Print allocate clause.
static DenseBoolArrayAttr makeDenseBoolArrayAttr(MLIRContext *ctx, const ArrayRef< bool > boolArray)
static std::string generateLoopNestingName(StringRef prefix, CanonicalLoopOp op)
Generate a name of a canonical loop nest of the format <prefix>(_r<idx>_s<idx>)*.
static ParseResult parseAffinityClause(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &iterated, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &affinityVars, SmallVectorImpl< Type > &iteratedTypes, SmallVectorImpl< Type > &affinityVarTypes)
static void printClauseWithRegionArgs(OpAsmPrinter &p, MLIRContext *ctx, StringRef clauseName, ValueRange argsSubrange, ValueRange operands, TypeRange types, ArrayAttr symbols=nullptr, DenseI64ArrayAttr mapIndices=nullptr, DenseBoolArrayAttr byref=nullptr, ReductionModifierAttr modifier=nullptr, UnitAttr needsBarrier=nullptr)
static void printSplitIteratedList(OpAsmPrinter &p, ValueRange iteratedVars, TypeRange iteratedTypes, ValueRange plainVars, TypeRange plainTypes, PrintPrefixFn &&printPrefixForPlain, PrintPrefixFn &&printPrefixForIterated)
static LogicalResult verifyDependVarList(Operation *op, std::optional< ArrayAttr > dependKinds, OperandRange dependVars, std::optional< ArrayAttr > iteratedKinds, OperandRange iteratedVars)
Verifies Depend clause.
static void printBlockArgClause(OpAsmPrinter &p, MLIRContext *ctx, StringRef clauseName, ValueRange argsSubrange, std::optional< MapPrintArgs > mapArgs)
static void printAffinityClause(OpAsmPrinter &p, Operation *op, ValueRange iterated, ValueRange affinityVars, TypeRange iteratedTypes, TypeRange affinityVarTypes)
static void printBlockArgRegion(OpAsmPrinter &p, Operation *op, Region ®ion, const AllRegionPrintArgs &args)
static ParseResult parseGranularityClause(OpAsmParser &parser, ClauseTypeAttr &prescriptiveness, std::optional< OpAsmParser::UnresolvedOperand > &operand, Type &operandType, std::optional< ClauseType >(*symbolizeClause)(StringRef), StringRef clauseName)
static void printIteratorHeader(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange lbs, ValueRange ubs, ValueRange steps, TypeRange, TypeRange, TypeRange)
static LogicalResult verifyDeclareTargetAttr(Operation *op, Attribute attr)
static ParseResult parseHeapAllocClause(OpAsmParser &parser, TypeAttr &inTypeAttr, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &typeparams, SmallVectorImpl< Type > &typeparamsTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &shape, SmallVectorImpl< Type > &shapeTypes)
operation ::= $in_type ( ( $typeparams ) )? ( , $shape )?
static void printInReductionClause(OpAsmPrinter &p, Operation *op, ValueRange inReductionVars, TypeRange inReductionTypes, DenseBoolArrayAttr inReductionByref, ArrayAttr inReductionSyms)
Prints an in_reduction clause for an operation that does not give its list items entry block argument...
static ParseResult parseIteratorHeader(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &lbs, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &ubs, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &steps, SmallVectorImpl< Type > &lbTypes, SmallVectorImpl< Type > &ubTypes, SmallVectorImpl< Type > &stepTypes)
static ParseResult parseBlockArgRegion(OpAsmParser &parser, Region ®ion, AllRegionParseArgs args)
static ParseResult parseLoopTransformClis(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &generateesOperands, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &applyeesOperands)
static ParseResult parseSynchronizationHint(OpAsmParser &parser, IntegerAttr &hintAttr)
Parses a Synchronization Hint clause.
static void printScheduleClause(OpAsmPrinter &p, Operation *op, ClauseScheduleKindAttr scheduleKind, ScheduleModifierAttr scheduleMod, UnitAttr scheduleSimd, Value scheduleChunk, Type scheduleChunkType)
Print schedule clause.
static void printCopyprivate(OpAsmPrinter &p, Operation *op, OperandRange copyprivateVars, TypeRange copyprivateTypes, std::optional< ArrayAttr > copyprivateSyms)
Print Copyprivate clause.
static ParseResult parseOrderClause(OpAsmParser &parser, ClauseOrderKindAttr &order, OrderModifierAttr &orderMod)
static bool mapTypeToBool(ClauseMapFlags value, ClauseMapFlags flag)
static void printAlignedClause(OpAsmPrinter &p, Operation *op, ValueRange alignedVars, TypeRange alignedTypes, std::optional< ArrayAttr > alignments)
Print Aligned Clause.
static bool targetInReductionCapturedBy(Value inReductionVar, Value mapVarPtr)
An omp.target in_reduction operand is captured by a map_entries entry when the entry's MapInfoOp var_...
static LogicalResult verifySynchronizationHint(Operation *op, uint64_t hint)
Verifies a synchronization hint clause.
static ParseResult parseUseDeviceAddrUseDevicePtrRegion(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &useDeviceAddrVars, SmallVectorImpl< Type > &useDeviceAddrTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &useDevicePtrVars, SmallVectorImpl< Type > &useDevicePtrTypes)
static ParseResult parseUniformClause(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &uniformVars, SmallVectorImpl< Type > &uniformTypes)
uniform ::= uniform ( uniform-list ) uniform-list := uniform-val (, uniform-val)* uniform-val := ssa-...
static void printInReductionPrivateReductionRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange inReductionVars, TypeRange inReductionTypes, DenseBoolArrayAttr inReductionByref, ArrayAttr inReductionSyms, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier, ReductionModifierAttr reductionMod, ValueRange reductionVars, TypeRange reductionTypes, DenseBoolArrayAttr reductionByref, ArrayAttr reductionSyms)
static void printInReductionPrivateRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange inReductionVars, TypeRange inReductionTypes, DenseBoolArrayAttr inReductionByref, ArrayAttr inReductionSyms, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier)
static LogicalResult verifyAllocateClause(Operation *op, ValueRange allocateVars, ValueRange allocatorVars, DenseI64ArrayAttr allocateAlignments, DenseI64ArrayAttr allocatePrivateIndices, ValueRange privateVars={}, ArrayAttr privateSyms=nullptr, bool requirePrivateIndices=false)
static void printSynchronizationHint(OpAsmPrinter &p, Operation *op, IntegerAttr hintAttr)
Prints a Synchronization Hint clause.
static void printGranularityClause(OpAsmPrinter &p, Operation *op, ClauseTypeAttr prescriptiveness, Value operand, mlir::Type operandType, StringRef(*stringifyClauseType)(ClauseType))
static ParseResult parseDependVarList(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &dependVars, SmallVectorImpl< Type > &dependTypes, ArrayAttr &dependKinds, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &iteratedVars, SmallVectorImpl< Type > &iteratedTypes, ArrayAttr &iteratedKinds)
depend-entry-list ::= depend-entry | depend-entry-list , depend-entry depend-entry ::= depend-kind ->...
static Operation * getParentInSameDialect(Operation *thisOp)
static void printUniformClause(OpAsmPrinter &p, Operation *op, ValueRange uniformVars, TypeRange uniformTypes)
Print Uniform Clauses.
static LogicalResult verifyCopyprivateVarList(Operation *op, OperandRange copyprivateVars, std::optional< ArrayAttr > copyprivateSyms)
Verifies CopyPrivate Clause.
static LogicalResult verifyAlignedClause(Operation *op, std::optional< ArrayAttr > alignments, OperandRange alignedVars)
static ParseResult parsePrivateRegion(OpAsmParser &parser, Region ®ion, llvm::SmallVectorImpl< OpAsmParser::UnresolvedOperand > &privateVars, llvm::SmallVectorImpl< Type > &privateTypes, ArrayAttr &privateSyms, UnitAttr &privateNeedsBarrier)
static void printNumTasksClause(OpAsmPrinter &p, Operation *op, ClauseNumTasksTypeAttr numTasksMod, Value numTasks, mlir::Type numTasksType)
static void printLoopTransformClis(OpAsmPrinter &p, TileOp op, OperandRange generatees, OperandRange applyees)
static ParseResult parseDynGroupprivateClause(OpAsmParser &parser, AccessGroupModifierAttr &accessGroupAttr, FallbackModifierAttr &fallbackAttr, std::optional< OpAsmParser::UnresolvedOperand > &dynGroupprivateSize, Type &sizeType)
static void printPrivateRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier)
static void printPrivateReductionRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier, ReductionModifierAttr reductionMod, ValueRange reductionVars, TypeRange reductionTypes, DenseBoolArrayAttr reductionByref, ArrayAttr reductionSyms)
static ParseResult parseSplitIteratedList(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &iteratedVars, SmallVectorImpl< Type > &iteratedTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &plainVars, SmallVectorImpl< Type > &plainTypes, ParsePrefixFn &&parsePrefix)
static void printTaskReductionRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange taskReductionVars, TypeRange taskReductionTypes, DenseBoolArrayAttr taskReductionByref, ArrayAttr taskReductionSyms)
static LogicalResult verifyMapInfoForMapClause(Operation *op, mlir::omp::MapInfoOp mapInfoOp, llvm::DenseSet< mlir::TypedValue< mlir::omp::PointerLikeType > > &updateToVars, llvm::DenseSet< mlir::TypedValue< mlir::omp::PointerLikeType > > &updateFromVars)
static LogicalResult verifyOrderedParent(Operation &op)
static void printOrderClause(OpAsmPrinter &p, Operation *op, ClauseOrderKindAttr order, OrderModifierAttr orderMod)
static ParseResult parseBlockArgClause(OpAsmParser &parser, llvm::SmallVectorImpl< OpAsmParser::Argument > &entryBlockArgs, StringRef keyword, std::optional< MapParseArgs > mapArgs)
static ParseResult parseClauseWithRegionArgs(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &operands, SmallVectorImpl< Type > &types, SmallVectorImpl< OpAsmParser::Argument > ®ionPrivateArgs, ArrayAttr *symbols=nullptr, DenseI64ArrayAttr *mapIndices=nullptr, DenseBoolArrayAttr *byref=nullptr, ReductionModifierAttr *modifier=nullptr, UnitAttr *needsBarrier=nullptr)
static LogicalResult verifyPrivateVarsMapping(TargetOp targetOp)
static ParseResult parseScheduleClause(OpAsmParser &parser, ClauseScheduleKindAttr &scheduleAttr, ScheduleModifierAttr &scheduleMod, UnitAttr &scheduleSimd, std::optional< OpAsmParser::UnresolvedOperand > &chunkSize, Type &chunkType)
schedule ::= schedule ( sched-list ) sched-list ::= sched-val | sched-val sched-list | sched-val ,...
static LogicalResult verifyDynGroupprivateClause(Operation *op, AccessGroupModifierAttr accessGroup, FallbackModifierAttr fallback, Value dynGroupprivateSize)
static LogicalResult verifyLinearModifiers(Operation *op, std::optional< ArrayAttr > linearModifiers, OperandRange linearVars, bool isDeclareSimd=false)
OpenMP 5.2, Section 5.4.6: "A linear-modifier may be specified as ref or uval only on a declare simd ...
static void printClauseAttr(OpAsmPrinter &p, Operation *op, ClauseAttr attr)
static ParseResult parseAllocateAndAllocator(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &allocateVars, SmallVectorImpl< Type > &allocateTypes, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &allocatorVars, SmallVectorImpl< Type > &allocatorTypes)
Parse an allocate clause with allocators and a list of operands with types.
static void printMembersIndex(OpAsmPrinter &p, MapInfoOp op, ArrayAttr membersIdx)
static void printCaptureType(OpAsmPrinter &p, Operation *op, VariableCaptureKindAttr mapCaptureType)
static LogicalResult verifyNumTeamsClause(Operation *op, Value numTeamsLower, OperandRange numTeamsUpperVars)
static bool opInGlobalImplicitParallelRegion(Operation *op)
static void printTargetOpRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange hasDeviceAddrVars, TypeRange hasDeviceAddrTypes, ValueRange hostEvalVars, TypeRange hostEvalTypes, ValueRange mapVars, TypeRange mapTypes, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier, DenseI64ArrayAttr privateMaps)
static void printUseDeviceAddrUseDevicePtrRegion(OpAsmPrinter &p, Operation *op, Region ®ion, ValueRange useDeviceAddrVars, TypeRange useDeviceAddrTypes, ValueRange useDevicePtrVars, TypeRange useDevicePtrTypes)
static LogicalResult verifyMapClause(Operation *op, OperandRange mapVars, OperandRange mapIterated)
static LogicalResult verifyPrivateVarList(OpType &op)
static ParseResult parseNumTasksClause(OpAsmParser &parser, ClauseNumTasksTypeAttr &numTasksMod, std::optional< OpAsmParser::UnresolvedOperand > &numTasks, Type &numTasksType)
LogicalResult verifyAlignment(Operation &op, std::optional< uint64_t > alignment)
Verifies align clause in allocate directive.
static ParseResult parseAlignedClause(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &alignedVars, SmallVectorImpl< Type > &alignedTypes, ArrayAttr &alignmentsAttr)
aligned ::= aligned ( aligned-list ) aligned-list := aligned-val | aligned-val aligned-list aligned-v...
static ParseResult parsePrivateReductionRegion(OpAsmParser &parser, Region ®ion, llvm::SmallVectorImpl< OpAsmParser::UnresolvedOperand > &privateVars, llvm::SmallVectorImpl< Type > &privateTypes, ArrayAttr &privateSyms, UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &reductionVars, SmallVectorImpl< Type > &reductionTypes, DenseBoolArrayAttr &reductionByref, ArrayAttr &reductionSyms)
static void printLinearClause(OpAsmPrinter &p, Operation *op, ValueRange linearVars, TypeRange linearTypes, ValueRange linearStepVars, TypeRange stepVarTypes, ArrayAttr linearModifiers)
Print Linear Clause.
static ParseResult parseInReductionPrivateReductionRegion(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &inReductionVars, SmallVectorImpl< Type > &inReductionTypes, DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms, llvm::SmallVectorImpl< OpAsmParser::UnresolvedOperand > &privateVars, llvm::SmallVectorImpl< Type > &privateTypes, ArrayAttr &privateSyms, UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &reductionVars, SmallVectorImpl< Type > &reductionTypes, DenseBoolArrayAttr &reductionByref, ArrayAttr &reductionSyms)
static LogicalResult checkApplyeesNesting(TileOp op)
Check properties of the loop nest consisting of the transformation's applyees:
static ParseResult parseCaptureType(OpAsmParser &parser, VariableCaptureKindAttr &mapCaptureType)
static ParseResult parseTaskReductionRegion(OpAsmParser &parser, Region ®ion, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &taskReductionVars, SmallVectorImpl< Type > &taskReductionTypes, DenseBoolArrayAttr &taskReductionByref, ArrayAttr &taskReductionSyms)
static ParseResult parseGrainsizeClause(OpAsmParser &parser, ClauseGrainsizeTypeAttr &grainsizeMod, std::optional< OpAsmParser::UnresolvedOperand > &grainsize, Type &grainsizeType)
static ParseResult parseCopyprivate(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > ©privateVars, SmallVectorImpl< Type > ©privateTypes, ArrayAttr ©privateSyms)
copyprivate-entry-list ::= copyprivate-entry | copyprivate-entry-list , copyprivate-entry copyprivate...
static ParseResult parseInReductionClause(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &inReductionVars, SmallVectorImpl< Type > &inReductionTypes, DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms)
Parses an in_reduction clause for an operation that does not give its list items entry block argument...
static LogicalResult verifyMapInfoDefinedArgs(Operation *op, StringRef clauseName, OperandRange vars)
static void printGrainsizeClause(OpAsmPrinter &p, Operation *op, ClauseGrainsizeTypeAttr grainsizeMod, Value grainsize, mlir::Type grainsizeType)
static ParseResult verifyScheduleModifiers(OpAsmParser &parser, SmallVectorImpl< SmallString< 12 > > &modifiers)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
static bool isUnique(It begin, It end)
static LogicalResult emit(SolverOp solver, const SMTEmissionOptions &options, mlir::raw_indented_ostream &stream)
Emit the SMT operations in the given 'solver' to the 'stream'.
static SmallVector< Value > getTileSizes(Location loc, x86::amx::TileType tType, RewriterBase &rewriter)
Maps the 2-dim vector shape to the two 16-bit tile sizes.
This base class exposes generic asm parser hooks, usable across the various derived parsers.
virtual ParseResult parseMinus()=0
Parse a '-' token.
@ Paren
Parens surrounding zero or more operands.
@ None
Zero or more operands with no delimiters.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseOptionalColon()=0
Parse a : token if present.
virtual ParseResult parseLSquare()=0
Parse a [ token.
virtual ParseResult parseRSquare()=0
Parse a ] token.
ParseResult parseInteger(IntT &result)
Parse an integer value from the stream.
virtual ParseResult parseOptionalArrow()=0
Parse a '->' token if present.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseOptionalLess()=0
Parse a '<' token if present.
virtual ParseResult parseArrow()=0
Parse a '->' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalLParen()=0
Parse a ( token if present.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
SuccessorRange getSuccessors()
Operation * getTerminator()
Get the terminator operation of this block.
bool mightHaveTerminator()
Return "true" if this block might have a terminator.
BlockArgListType getArguments()
IntegerAttr getI64IntegerAttr(int64_t value)
IntegerType getIntegerType(unsigned width)
MLIRContext * getContext() const
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Diagnostic & append(Arg1 &&arg1, Arg2 &&arg2, Args &&...args)
Append arguments to the diagnostic.
Diagnostic & appendOp(Operation &op, const OpPrintingFlags &flags)
Append an operation with the given printing flags.
A class for computing basic dominance information.
bool dominates(Operation *a, Operation *b) const
Return true if operation A dominates operation B, i.e.
This class represents a diagnostic that is inflight and set to be reported.
Diagnostic & attachNote(std::optional< Location > noteLoc=std::nullopt)
Attaches a note to this diagnostic.
MLIRContext is the top-level object for a collection of MLIR operations.
NamedAttribute represents a combination of a name and an Attribute value.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgument(Argument &result, bool allowType=false, bool allowAttrs=false)=0
Parse a single argument with the following syntax:
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
virtual void printRegionArgument(BlockArgument arg, ArrayRef< NamedAttribute > argAttrs={}, bool omitType=false)=0
Print a block argument in the usual format of: ssaName : type {attr1=42} loc("here") where location p...
virtual void printOperand(Value value)=0
Print implementations for various things an operation contains.
This class helps build Operations.
This class represents an operand of an operation.
Set of flags used to control the behavior of the various IR print methods (e.g.
This class provides the API for ops that are known to be isolated from above.
This class provides the API for ops that are known to be terminators.
This class indicates that the regions associated with this op don't have terminators.
This class implements the operand iterators for the Operation class.
type_range getType() const
Operation is the basic unit of execution within MLIR.
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Block * getBlock()
Returns the operation block that contains this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Location getLoc()
The source location the operation was defined or derived from.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
user_range getUsers()
Returns a range of all users.
Region * getParentRegion()
Returns the region to which the instruction belongs.
MLIRContext * getContext()
Return the context this operation is associated with.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockArgListType getArguments()
OpIterator op_begin()
Return iterators that walk the operations nested directly within this region.
bool isAncestor(Region *other)
Return true if this region is ancestor of the other region.
iterator_range< OpIterator > getOps()
unsigned getNumArguments()
Location getLoc()
Return a location for this region.
BlockArgument getArgument(unsigned i)
Operation * getParentOp()
Return the parent operation this region is attached to.
BlockListType & getBlocks()
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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 Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
type_range getType() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Type getType() const
Return the type of this value.
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
A utility result that is used to signal how to proceed with an ongoing walk:
static WalkResult advance()
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< bool > content)
ArrayRef< T > asArrayRef() const
bool isReachableFromEntry(Block *a) const
Return true if the specified block is reachable from the entry block of its region.
Operation * getOwner() const
Return the owner of this operand.
TargetEnterDataOperands TargetEnterExitUpdateDataOperands
omp.target_enter_data, omp.target_exit_data and omp.target_update take the same clauses,...
std::tuple< NewCliOp, OpOperand *, OpOperand * > decodeCli(mlir::Value cli)
Find the omp.new_cli, generator, and consumer of a canonical loop info.
ClauseProcBindKind convertProcBindKind(llvm::omp::ProcBindKind kind)
Convert a proc_bind kind from the LLVM frontend enum to the corresponding OpenMP dialect enum.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
function_ref< void(Value, StringRef)> OpAsmSetValueNameFn
A functor used to set the name of the start of a result group of an operation.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isPure(Operation *op)
Returns true if the given operation is pure, i.e., is speculatable that does not touch memory.
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
llvm::TypeSwitch< T, ResultT > TypeSwitch
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
detail::DenseArrayAttrImpl< bool > DenseBoolArrayAttr
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
function_ref< void(Block *, StringRef)> OpAsmSetBlockNameFn
A functor used to set the name of blocks in regions directly nested under an operation.
This is the representation of an operand reference.
This class provides APIs and verifiers for ops with regions having a single block.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
T & getOrAddProperties()
Get (or create) the properties of the provided type to be set on the operation on creation.
void addOperands(ValueRange newOperands)
void addAttributes(ArrayRef< NamedAttribute > newAttributes)
Add an array of named attributes.
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addTypes(ArrayRef< Type > newTypes)
Region * addRegion()
Create a region that should be attached to the operation.
Extended TargetOperands with kernel_type attribute.
TargetExecModeAttr kernelType
Kernel execution mode for the target region.