MLIR 24.0.0git
OpenMPDialect.cpp
Go to the documentation of this file.
1//===- OpenMPDialect.cpp - MLIR Dialect for OpenMP implementation ---------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements the OpenMP dialect and its operations.
10//
11//===----------------------------------------------------------------------===//
12
18#include "mlir/IR/Attributes.h"
21#include "mlir/IR/Matchers.h"
24#include "mlir/IR/SymbolTable.h"
27
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"
38#include <cstddef>
39#include <iterator>
40#include <optional>
41#include <variant>
42
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"
47
48using namespace mlir;
49using namespace mlir::omp;
50
53 return attrs.empty() ? nullptr : ArrayAttr::get(context, attrs);
54}
55
58 return boolArray.empty() ? nullptr : DenseBoolArrayAttr::get(ctx, boolArray);
59}
60
63 return intArray.empty() ? nullptr : DenseI64ArrayAttr::get(ctx, intArray);
64}
65
66namespace {
67struct MemRefPointerLikeModel
68 : public PointerLikeType::ExternalModel<MemRefPointerLikeModel,
69 MemRefType> {
70 Type getElementType(Type pointer) const {
71 return llvm::cast<MemRefType>(pointer).getElementType();
72 }
73};
74
75struct LLVMPointerPointerLikeModel
76 : public PointerLikeType::ExternalModel<LLVMPointerPointerLikeModel,
77 LLVM::LLVMPointerType> {
78 Type getElementType(Type pointer) const { return Type(); }
79};
80} // namespace
81
82/// Generate a name of a canonical loop nest of the format
83/// `<prefix>(_r<idx>_s<idx>)*`. Hereby, `_r<idx>` identifies the region
84/// argument index of an operation that has multiple regions, if the operation
85/// has multiple regions.
86/// `_s<idx>` identifies the position of an operation within a region, where
87/// only operations that may potentially contain loops ("container operations"
88/// i.e. have region arguments) are counted. Again, it is omitted if there is
89/// only one such operation in a region. If there are canonical loops nested
90/// inside each other, also may also use the format `_d<num>` where <num> is the
91/// nesting depth of the loop.
92///
93/// The generated name is a best-effort to make canonical loop unique within an
94/// SSA namespace. This also means that regions with IsolatedFromAbove property
95/// do not consider any parents or siblings.
96static std::string generateLoopNestingName(StringRef prefix,
97 CanonicalLoopOp op) {
98 struct Component {
99 /// If true, this component describes a region operand of an operation (the
100 /// operand's owner) If false, this component describes an operation located
101 /// in a parent region
102 bool isRegionArgOfOp;
103 bool skip = false;
104 bool isUnique = false;
105
106 size_t idx;
107 Operation *op;
108 Region *parentRegion;
109 size_t loopDepth;
110
111 Operation *&getOwnerOp() {
112 assert(isRegionArgOfOp && "Must describe a region operand");
113 return op;
114 }
115 size_t &getArgIdx() {
116 assert(isRegionArgOfOp && "Must describe a region operand");
117 return idx;
118 }
119
120 Operation *&getContainerOp() {
121 assert(!isRegionArgOfOp && "Must describe a operation of a region");
122 return op;
123 }
124 size_t &getOpPos() {
125 assert(!isRegionArgOfOp && "Must describe a operation of a region");
126 return idx;
127 }
128 bool isLoopOp() const {
129 assert(!isRegionArgOfOp && "Must describe a operation of a region");
130 return isa<CanonicalLoopOp>(op);
131 }
132 Region *&getParentRegion() {
133 assert(!isRegionArgOfOp && "Must describe a operation of a region");
134 return parentRegion;
135 }
136 size_t &getLoopDepth() {
137 assert(!isRegionArgOfOp && "Must describe a operation of a region");
138 return loopDepth;
139 }
140
141 void skipIf(bool v = true) { skip = skip || v; }
142 };
143
144 // List of ancestors, from inner to outer.
145 // Alternates between
146 // * region argument of an operation
147 // * operation within a region
148 SmallVector<Component> components;
149
150 // Gather a list of parent regions and operations, and the position within
151 // their parent
152 Operation *o = op.getOperation();
153 while (o) {
154 // Operation within a region
155 Region *r = o->getParentRegion();
156 if (!r)
157 break;
158
159 llvm::ReversePostOrderTraversal<Block *> traversal(&r->getBlocks().front());
160 size_t idx = 0;
161 bool found = false;
162 size_t sequentialIdx = -1;
163 bool isOnlyContainerOp = true;
164 for (Block *b : traversal) {
165 for (Operation &op : *b) {
166 if (&op == o && !found) {
167 sequentialIdx = idx;
168 found = true;
169 }
170 if (op.getNumRegions()) {
171 idx += 1;
172 if (idx > 1)
173 isOnlyContainerOp = false;
174 }
175 if (found && !isOnlyContainerOp)
176 break;
177 }
178 }
179
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;
186
187 Operation *parent = r->getParentOp();
188
189 // Region argument of an operation
190 Component &regionArgOfOperation = components.emplace_back();
191 regionArgOfOperation.isRegionArgOfOp = true;
192 regionArgOfOperation.isUnique = true;
193 regionArgOfOperation.getArgIdx() = 0;
194 regionArgOfOperation.getOwnerOp() = parent;
195
196 // The IsolatedFromAbove trait of the parent operation implies that each
197 // individual region argument has its own separate namespace, so no
198 // ambiguity.
199 if (!parent || parent->hasTrait<mlir::OpTrait::IsIsolatedFromAbove>())
200 break;
201
202 // Component only needed if operation has multiple region operands. Region
203 // arguments may be optional, but we currently do not consider this.
204 if (parent->getRegions().size() > 1) {
205 auto getRegionIndex = [](Operation *o, Region *r) {
206 for (auto [idx, region] : llvm::enumerate(o->getRegions())) {
207 if (&region == r)
208 return idx;
209 }
210 llvm_unreachable("Region not child of its parent operation");
211 };
212 regionArgOfOperation.isUnique = false;
213 regionArgOfOperation.getArgIdx() = getRegionIndex(parent, r);
214 }
215
216 // next parent
217 o = parent;
218 }
219
220 // Determine whether a region-argument component is not needed
221 for (Component &c : components)
222 c.skipIf(c.isRegionArgOfOp && c.isUnique);
223
224 // Find runs of nested loops and determine each loop's depth in the loop nest
225 size_t numSurroundingLoops = 0;
226 for (Component &c : llvm::reverse(components)) {
227 if (c.skip)
228 continue;
229
230 // non-skipped multi-argument operands interrupt the loop nest
231 if (c.isRegionArgOfOp) {
232 numSurroundingLoops = 0;
233 continue;
234 }
235
236 // Multiple loops in a region means each of them is the outermost loop of a
237 // new loop nest
238 if (!c.isUnique)
239 numSurroundingLoops = 0;
240
241 c.getLoopDepth() = numSurroundingLoops;
242
243 // Next loop is surrounded by one more loop
244 if (isa<CanonicalLoopOp>(c.getContainerOp()))
245 numSurroundingLoops += 1;
246 }
247
248 // In loop nests, skip all but the innermost loop that contains the depth
249 // number
250 bool isLoopNest = false;
251 for (Component &c : components) {
252 if (c.skip || c.isRegionArgOfOp)
253 continue;
254
255 if (!isLoopNest && c.getLoopDepth() >= 1) {
256 // Innermost loop of a loop nest of at least two loops
257 isLoopNest = true;
258 } else if (isLoopNest) {
259 // Non-innermost loop of a loop nest
260 c.skipIf(c.isUnique);
261
262 // If there is no surrounding loop left, this must have been the outermost
263 // loop; leave loop-nest mode for the next iteration
264 if (c.getLoopDepth() == 0)
265 isLoopNest = false;
266 }
267 }
268
269 // Skip non-loop unambiguous regions (but they should interrupt loop nests, so
270 // we mark them as skipped only after computing loop nests)
271 for (Component &c : components)
272 c.skipIf(!c.isRegionArgOfOp && c.isUnique &&
273 !isa<CanonicalLoopOp>(c.getContainerOp()));
274
275 // Components can be skipped if they are already disambiguated by their parent
276 // (or does not have a parent)
277 bool newRegion = true;
278 for (Component &c : llvm::reverse(components)) {
279 c.skipIf(newRegion && c.isUnique);
280
281 // non-skipped components disambiguate unique children
282 if (!c.skip)
283 newRegion = true;
284
285 // ...except canonical loops that need a suffix for each nest
286 if (!c.isRegionArgOfOp && c.getContainerOp())
287 newRegion = false;
288 }
289
290 // Compile the nesting name string
291 SmallString<64> Name{prefix};
292 llvm::raw_svector_ostream NameOS(Name);
293 for (auto &c : llvm::reverse(components)) {
294 if (c.skip)
295 continue;
296
297 if (c.isRegionArgOfOp)
298 NameOS << "_r" << c.getArgIdx();
299 else if (c.getLoopDepth() >= 1)
300 NameOS << "_d" << c.getLoopDepth();
301 else
302 NameOS << "_s" << c.getOpPos();
303 }
304
305 return NameOS.str().str();
306}
307
308void OpenMPDialect::initialize() {
309 addOperations<
310#define GET_OP_LIST
311#include "mlir/Dialect/OpenMP/OpenMPOps.cpp.inc"
312 >();
313 addAttributes<
314#define GET_ATTRDEF_LIST
315#include "mlir/Dialect/OpenMP/OpenMPOpsAttributes.cpp.inc"
316 >();
317 addTypes<
318#define GET_TYPEDEF_LIST
319#include "mlir/Dialect/OpenMP/OpenMPOpsTypes.cpp.inc"
320 >();
321
322 declarePromisedInterface<ConvertToLLVMPatternInterface, OpenMPDialect>();
323
324 MemRefType::attachInterface<MemRefPointerLikeModel>(*getContext());
325 LLVM::LLVMPointerType::attachInterface<LLVMPointerPointerLikeModel>(
326 *getContext());
327
328 // Attach default offload module interface to module op to access
329 // offload functionality through
330 mlir::ModuleOp::attachInterface<mlir::omp::OffloadModuleDefaultModel>(
331 *getContext());
332
333 // Attach default declare target interfaces to operations which can be marked
334 // as declare target (Global Operations and Functions/Subroutines in dialects
335 // that Fortran (or other languages that lower to MLIR) translates too
336 mlir::LLVM::GlobalOp::attachInterface<
338 *getContext());
339 mlir::LLVM::LLVMFuncOp::attachInterface<
341 *getContext());
342 mlir::func::FuncOp::attachInterface<
344}
345
346//===----------------------------------------------------------------------===//
347// Dialect operation attribute verification
348//===----------------------------------------------------------------------===//
349
350static LogicalResult verifyDeclareTargetAttr(Operation *op, Attribute attr) {
351 if (!isa<DeclareTargetInterface>(op))
352 return op->emitError() << "omp.declare_target can only be applied to "
353 "DeclareTargetInterface ops";
354
355 auto declareTargetAttr = dyn_cast<DeclareTargetAttr>(attr);
356 if (!declareTargetAttr)
357 return op->emitError()
358 << "omp.declare_target must be an #omp.declaretarget attribute";
359
360 if (isa<mlir::FunctionOpInterface>(op)) {
361 if (declareTargetAttr.getAutomap())
362 return op->emitOpError()
363 << "omp.declare_target 'automap' is not valid on functions";
364
365 // TODO: Disallow the `local` clause (OpenMP 6.0).
366 if (declareTargetAttr.getCaptureClause() ==
367 mlir::omp::DeclareTargetCaptureClause::link)
368 return op->emitOpError()
369 << "omp.declare_target 'link' is not valid on functions";
370 } else {
371 // TODO: Disallow the `indirect` clause (OpenMP 5.1).
372 if (declareTargetAttr.getImplicit())
373 return op->emitOpError()
374 << "omp.declare_target 'implicit' is only valid on functions";
375 }
376 return success();
377}
378
379LogicalResult
380OpenMPDialect::verifyOperationAttribute(Operation *op,
381 NamedAttribute attribute) {
382 if (attribute.getName() == "omp.declare_target")
383 return verifyDeclareTargetAttr(op, attribute.getValue());
384
385 return success();
386}
387
388//===----------------------------------------------------------------------===//
389// Parser and printer for Allocate Clause
390//===----------------------------------------------------------------------===//
391
392/// Parse an allocate clause with allocators and a list of operands with types.
393///
394/// allocate-operand-list :: = allocate-operand |
395/// allocator-operand `,` allocate-operand-list
396/// allocate-operand :: = ssa-id-and-type -> ssa-id-and-type
397/// ssa-id-and-type ::= ssa-id `:` type
398static ParseResult parseAllocateAndAllocator(
399 OpAsmParser &parser,
401 SmallVectorImpl<Type> &allocateTypes,
403 SmallVectorImpl<Type> &allocatorTypes) {
404
405 return parser.parseCommaSeparatedList([&]() {
407 Type type;
408 if (parser.parseOperand(operand) || parser.parseColonType(type))
409 return failure();
410 allocatorVars.push_back(operand);
411 allocatorTypes.push_back(type);
412 if (parser.parseArrow())
413 return failure();
414 if (parser.parseOperand(operand) || parser.parseColonType(type))
415 return failure();
416
417 allocateVars.push_back(operand);
418 allocateTypes.push_back(type);
419 return success();
420 });
421}
422
423/// Print allocate clause
425 OperandRange allocateVars,
426 TypeRange allocateTypes,
427 OperandRange allocatorVars,
428 TypeRange allocatorTypes) {
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;
433 }
434}
435
436//===----------------------------------------------------------------------===//
437// Parser and printer for a clause attribute (StringEnumAttr)
438//===----------------------------------------------------------------------===//
439
440template <typename ClauseAttr>
441static ParseResult parseClauseAttr(AsmParser &parser, ClauseAttr &attr) {
442 using ClauseT = decltype(std::declval<ClauseAttr>().getValue());
443 StringRef enumStr;
444 SMLoc loc = parser.getCurrentLocation();
445 if (parser.parseKeyword(&enumStr))
446 return failure();
447 if (std::optional<ClauseT> enumValue = symbolizeEnum<ClauseT>(enumStr)) {
448 attr = ClauseAttr::get(parser.getContext(), *enumValue);
449 return success();
450 }
451 return parser.emitError(loc, "invalid clause value: '") << enumStr << "'";
452}
453
454template <typename ClauseAttr>
455static void printClauseAttr(OpAsmPrinter &p, Operation *op, ClauseAttr attr) {
456 p << stringifyEnum(attr.getValue());
457}
458
459//===----------------------------------------------------------------------===//
460// Parser and printer for Linear Clause
461//===----------------------------------------------------------------------===//
462
463/// linear ::= `linear` `(` linear-list `)`
464/// linear-list := linear-val | linear-val linear-list
465/// linear-val := ssa-id-and-type `=` ssa-id-and-type
466/// | `val` `(` ssa-id-and-type `=` ssa-id-and-type `)`
467/// | `ref` `(` ssa-id-and-type `=` ssa-id-and-type `)`
468/// | `uval` `(` ssa-id-and-type `=` ssa-id-and-type `)`
469static ParseResult parseLinearClause(
470 OpAsmParser &parser,
472 SmallVectorImpl<Type> &linearTypes,
474 SmallVectorImpl<Type> &linearStepTypes, ArrayAttr &linearModifiers) {
475 SmallVector<Attribute> modifiers;
476 auto result = parser.parseCommaSeparatedList([&]() {
478 Type type, stepType;
480
481 std::optional<omp::LinearModifier> linearModifier;
482 if (succeeded(parser.parseOptionalKeyword("val"))) {
483 linearModifier = omp::LinearModifier::val;
484 } else if (succeeded(parser.parseOptionalKeyword("ref"))) {
485 linearModifier = omp::LinearModifier::ref;
486 } else if (succeeded(parser.parseOptionalKeyword("uval"))) {
487 linearModifier = omp::LinearModifier::uval;
488 }
489
490 bool hasLinearModifierParens = linearModifier.has_value();
491 if (hasLinearModifierParens && parser.parseLParen())
492 return failure();
493
494 if (parser.parseOperand(var) || parser.parseColonType(type) ||
495 parser.parseEqual() || parser.parseOperand(stepVar) ||
496 parser.parseColonType(stepType))
497 return failure();
498
499 if (hasLinearModifierParens && parser.parseRParen())
500 return failure();
501
502 linearVars.push_back(var);
503 linearTypes.push_back(type);
504 linearStepVars.push_back(stepVar);
505 linearStepTypes.push_back(stepType);
506 if (linearModifier) {
507 modifiers.push_back(
508 omp::LinearModifierAttr::get(parser.getContext(), *linearModifier));
509 } else {
510 modifiers.push_back(UnitAttr::get(parser.getContext()));
511 }
512 return success();
513 });
514 if (failed(result))
515 return failure();
516 linearModifiers = ArrayAttr::get(parser.getContext(), modifiers);
517 return success();
518}
519
520/// Print Linear Clause
522 ValueRange linearVars, TypeRange linearTypes,
523 ValueRange linearStepVars, TypeRange stepVarTypes,
524 ArrayAttr linearModifiers) {
525 size_t linearVarsSize = linearVars.size();
526 for (unsigned i = 0; i < linearVarsSize; ++i) {
527 if (i != 0)
528 p << ", ";
529 // Print modifier keyword wrapper if present.
530 Attribute modAttr = linearModifiers ? linearModifiers[i] : nullptr;
531 auto mod = modAttr ? dyn_cast<omp::LinearModifierAttr>(modAttr) : nullptr;
532 if (mod) {
533 p << omp::stringifyLinearModifier(mod.getValue()) << "(";
534 }
535 p << linearVars[i] << " : " << linearTypes[i];
536 p << " = " << linearStepVars[i] << " : " << stepVarTypes[i];
537 if (mod)
538 p << ")";
539 }
540}
541
542//===----------------------------------------------------------------------===//
543// Verifier for Linear modifier
544//===----------------------------------------------------------------------===//
545
546/// OpenMP 5.2, Section 5.4.6: "A linear-modifier may be specified as ref or
547/// uval only on a declare simd directive."
548/// Also verifies that modifier count matches variable count.
549static LogicalResult
550verifyLinearModifiers(Operation *op, std::optional<ArrayAttr> linearModifiers,
551 OperandRange linearVars, bool isDeclareSimd = false) {
552 if (!linearModifiers)
553 return success();
554 if (linearModifiers->size() != linearVars.size())
555 return op->emitOpError()
556 << "expected as many linear modifiers as linear variables";
557 if (!isDeclareSimd) {
558 for (Attribute attr : *linearModifiers) {
559 if (!attr)
560 continue;
561 auto modAttr = dyn_cast<omp::LinearModifierAttr>(attr);
562 if (!modAttr)
563 continue;
564 omp::LinearModifier mod = modAttr.getValue();
565 if (mod == omp::LinearModifier::ref || mod == omp::LinearModifier::uval)
566 return op->emitOpError()
567 << "linear modifier '" << omp::stringifyLinearModifier(mod)
568 << "' may only be specified on a declare simd directive";
569 }
570 }
571 return success();
572}
573
574//===----------------------------------------------------------------------===//
575// Verifier for Nontemporal Clause
576//===----------------------------------------------------------------------===//
577
578static LogicalResult verifyNontemporalClause(Operation *op,
579 OperandRange nontemporalVars) {
580
581 // Check if each var is unique - OpenMP 5.0 -> 2.9.3.1 section
582 DenseSet<Value> nontemporalItems;
583 for (const auto &it : nontemporalVars)
584 if (!nontemporalItems.insert(it).second)
585 return op->emitOpError() << "nontemporal variable used more than once";
586
587 return success();
588}
589
590//===----------------------------------------------------------------------===//
591// Parser, verifier and printer for Aligned Clause
592//===----------------------------------------------------------------------===//
593static LogicalResult verifyAlignedClause(Operation *op,
594 std::optional<ArrayAttr> alignments,
595 OperandRange alignedVars) {
596 // Check if number of alignment values equals to number of aligned variables
597 if (!alignedVars.empty()) {
598 if (!alignments || alignments->size() != alignedVars.size())
599 return op->emitOpError()
600 << "expected as many alignment values as aligned variables";
601 } else {
602 if (alignments)
603 return op->emitOpError() << "unexpected alignment values attribute";
604 return success();
605 }
606
607 // Check if each var is aligned only once - OpenMP 4.5 -> 2.8.1 section
608 DenseSet<Value> alignedItems;
609 for (auto it : alignedVars)
610 if (!alignedItems.insert(it).second)
611 return op->emitOpError() << "aligned variable used more than once";
612
613 if (!alignments)
614 return success();
615
616 // Check if all alignment values are positive - OpenMP 4.5 -> 2.8.1 section
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";
621 } else {
622 return op->emitOpError() << "expected integer alignment";
623 }
624 }
625
626 return success();
627}
628
629/// aligned ::= `aligned` `(` aligned-list `)`
630/// aligned-list := aligned-val | aligned-val aligned-list
631/// aligned-val := ssa-id-and-type `->` alignment
632static ParseResult
635 SmallVectorImpl<Type> &alignedTypes,
636 ArrayAttr &alignmentsAttr) {
637 SmallVector<Attribute> alignmentVec;
638 if (failed(parser.parseCommaSeparatedList([&]() {
639 if (parser.parseOperand(alignedVars.emplace_back()) ||
640 parser.parseColonType(alignedTypes.emplace_back()) ||
641 parser.parseArrow() ||
642 parser.parseAttribute(alignmentVec.emplace_back())) {
643 return failure();
644 }
645 return success();
646 })))
647 return failure();
648 SmallVector<Attribute> alignments(alignmentVec.begin(), alignmentVec.end());
649 alignmentsAttr = ArrayAttr::get(parser.getContext(), alignments);
650 return success();
651}
652
653/// Print Aligned Clause
655 ValueRange alignedVars, TypeRange alignedTypes,
656 std::optional<ArrayAttr> alignments) {
657 for (unsigned i = 0; i < alignedVars.size(); ++i) {
658 if (i != 0)
659 p << ", ";
660 p << alignedVars[i] << " : " << alignedVars[i].getType();
661 p << " -> " << (*alignments)[i];
662 }
663}
664
665static LogicalResult verifyAllocateClause(
666 Operation *op, ValueRange allocateVars, ValueRange allocatorVars,
667 DenseI64ArrayAttr allocateAlignments,
668 DenseI64ArrayAttr allocatePrivateIndices, ValueRange privateVars = {},
669 ArrayAttr privateSyms = nullptr, bool requirePrivateIndices = false) {
670 if (allocateVars.size() != allocatorVars.size())
671 return op->emitError(
672 "expected equal sizes for allocate and allocator variables");
673
674 if (allocateVars.empty()) {
675 if (allocateAlignments)
676 return op->emitError(
677 "unexpected allocate alignments without allocate variables");
678 if (allocatePrivateIndices)
679 return op->emitError(
680 "unexpected allocate private indices without allocate variables");
681 return success();
682 }
683
684 if (allocateAlignments) {
685 ArrayRef<int64_t> alignments = allocateAlignments.asArrayRef();
686 if (alignments.size() != allocateVars.size())
687 return op->emitError(
688 "expected as many allocate alignments as allocate variables");
689 for (int64_t alignment : alignments) {
690 if (alignment < 0)
691 return op->emitError("expected non-negative allocate alignments");
692 if (alignment != 0 && (alignment & (alignment - 1)) != 0)
693 return op->emitError(
694 "expected positive allocate alignments to be powers of two");
695 }
696 }
697
698 if (!allocatePrivateIndices) {
699 if (requirePrivateIndices)
700 return op->emitError(
701 "expected an allocate private index for each allocate variable");
702 return success();
703 }
704
705 ArrayRef<int64_t> indices = allocatePrivateIndices.asArrayRef();
706 if (indices.size() != allocateVars.size())
707 return op->emitError(
708 "expected as many allocate private indices as allocate variables");
709
710 DenseSet<int64_t> usedPrivateSlots;
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)
717 return op->emitError(
718 "allocate private index refers to a private variable more than once");
719
720 Value privateVar = privateVars[privateIndex];
721 if (allocateVar.getType() != privateVar.getType())
722 return op->emitError()
723 << "type mismatch between allocate variable and private variable "
724 "at index "
725 << privateIndex;
726 if (allocateVar != privateVar)
727 return op->emitError()
728 << "allocate variable does not match private variable at index "
729 << privateIndex;
730
731 if (!privateSyms ||
732 static_cast<uint64_t>(privateIndex) >= privateSyms.size())
733 return op->emitError(
734 "allocate private index does not have a privatizer symbol");
735
736 auto privateSym = dyn_cast<SymbolRefAttr>(privateSyms[privateIndex]);
737 if (!privateSym)
738 return op->emitError(
739 "allocate private index does not reference a privatizer symbol");
740 PrivateClauseOp privatizer =
742 if (!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)
747 return op->emitError(
748 "allocate private index must refer to private or firstprivate "
749 "storage");
750 }
751
752 return success();
753}
754
755//===----------------------------------------------------------------------===//
756// Parser, printer and verifier for Schedule Clause
757//===----------------------------------------------------------------------===//
758
759static ParseResult
761 SmallVectorImpl<SmallString<12>> &modifiers) {
762 if (modifiers.size() > 2)
763 return parser.emitError(parser.getNameLoc()) << " unexpected modifier(s)";
764 for (const auto &mod : modifiers) {
765 // Translate the string. If it has no value, then it was not a valid
766 // modifier!
767 auto symbol = symbolizeScheduleModifier(mod);
768 if (!symbol)
769 return parser.emitError(parser.getNameLoc())
770 << " unknown modifier type: " << mod;
771 }
772
773 // If we have one modifier that is "simd", then stick a "none" modiifer in
774 // index 0.
775 if (modifiers.size() == 1) {
776 if (symbolizeScheduleModifier(modifiers[0]) == ScheduleModifier::simd) {
777 modifiers.push_back(modifiers[0]);
778 modifiers[0] = stringifyScheduleModifier(ScheduleModifier::none);
779 }
780 } else if (modifiers.size() == 2) {
781 // If there are two modifier:
782 // First modifier should not be simd, second one should be simd
783 if (symbolizeScheduleModifier(modifiers[0]) == ScheduleModifier::simd ||
784 symbolizeScheduleModifier(modifiers[1]) != ScheduleModifier::simd)
785 return parser.emitError(parser.getNameLoc())
786 << " incorrect modifier order";
787 }
788 return success();
789}
790
791/// schedule ::= `schedule` `(` sched-list `)`
792/// sched-list ::= sched-val | sched-val sched-list |
793/// sched-val `,` sched-modifier
794/// sched-val ::= sched-with-chunk | sched-wo-chunk
795/// sched-with-chunk ::= sched-with-chunk-types (`=` ssa-id-and-type)?
796/// sched-with-chunk-types ::= `static` | `dynamic` | `guided`
797/// sched-wo-chunk ::= `auto` | `runtime`
798/// sched-modifier ::= sched-mod-val | sched-mod-val `,` sched-mod-val
799/// sched-mod-val ::= `monotonic` | `nonmonotonic` | `simd` | `none`
800static ParseResult
801parseScheduleClause(OpAsmParser &parser, ClauseScheduleKindAttr &scheduleAttr,
802 ScheduleModifierAttr &scheduleMod, UnitAttr &scheduleSimd,
803 std::optional<OpAsmParser::UnresolvedOperand> &chunkSize,
804 Type &chunkType) {
805 StringRef keyword;
806 if (parser.parseKeyword(&keyword))
807 return failure();
808 std::optional<mlir::omp::ClauseScheduleKind> schedule =
809 symbolizeClauseScheduleKind(keyword);
810 if (!schedule)
811 return parser.emitError(parser.getNameLoc()) << " expected schedule kind";
812
813 scheduleAttr = ClauseScheduleKindAttr::get(parser.getContext(), *schedule);
814 switch (*schedule) {
815 case ClauseScheduleKind::Static:
816 case ClauseScheduleKind::Dynamic:
817 case ClauseScheduleKind::Guided:
818 if (succeeded(parser.parseOptionalEqual())) {
819 chunkSize = OpAsmParser::UnresolvedOperand{};
820 if (parser.parseOperand(*chunkSize) || parser.parseColonType(chunkType))
821 return failure();
822 } else {
823 chunkSize = std::nullopt;
824 }
825 break;
826 case ClauseScheduleKind::Auto:
827 case ClauseScheduleKind::Runtime:
828 case ClauseScheduleKind::Distribute:
829 chunkSize = std::nullopt;
830 }
831
832 // If there is a comma, we have one or more modifiers..
834 while (succeeded(parser.parseOptionalComma())) {
835 StringRef mod;
836 if (parser.parseKeyword(&mod))
837 return failure();
838 modifiers.push_back(mod);
839 }
840
841 if (verifyScheduleModifiers(parser, modifiers))
842 return failure();
843
844 if (!modifiers.empty()) {
845 SMLoc loc = parser.getCurrentLocation();
846 if (std::optional<ScheduleModifier> mod =
847 symbolizeScheduleModifier(modifiers[0])) {
848 scheduleMod = ScheduleModifierAttr::get(parser.getContext(), *mod);
849 } else {
850 return parser.emitError(loc, "invalid schedule modifier");
851 }
852 // Only SIMD attribute is allowed here!
853 if (modifiers.size() > 1) {
854 assert(symbolizeScheduleModifier(modifiers[1]) == ScheduleModifier::simd);
855 scheduleSimd = UnitAttr::get(parser.getBuilder().getContext());
856 }
857 }
858
859 return success();
860}
861
862/// Print schedule clause
864 ClauseScheduleKindAttr scheduleKind,
865 ScheduleModifierAttr scheduleMod,
866 UnitAttr scheduleSimd, Value scheduleChunk,
867 Type scheduleChunkType) {
868 p << stringifyClauseScheduleKind(scheduleKind.getValue());
869 if (scheduleChunk)
870 p << " = " << scheduleChunk << " : " << scheduleChunk.getType();
871 if (scheduleMod)
872 p << ", " << stringifyScheduleModifier(scheduleMod.getValue());
873 if (scheduleSimd)
874 p << ", simd";
875}
876
877//===----------------------------------------------------------------------===//
878// Parser and printer for Order Clause
879//===----------------------------------------------------------------------===//
880
881// order ::= `order` `(` [order-modifier ':'] concurrent `)`
882// order-modifier ::= reproducible | unconstrained
883static ParseResult parseOrderClause(OpAsmParser &parser,
884 ClauseOrderKindAttr &order,
885 OrderModifierAttr &orderMod) {
886 StringRef enumStr;
887 SMLoc loc = parser.getCurrentLocation();
888 if (parser.parseKeyword(&enumStr))
889 return failure();
890 if (std::optional<OrderModifier> enumValue =
891 symbolizeOrderModifier(enumStr)) {
892 orderMod = OrderModifierAttr::get(parser.getContext(), *enumValue);
893 if (parser.parseOptionalColon())
894 return failure();
895 loc = parser.getCurrentLocation();
896 if (parser.parseKeyword(&enumStr))
897 return failure();
898 }
899 if (std::optional<ClauseOrderKind> enumValue =
900 symbolizeClauseOrderKind(enumStr)) {
901 order = ClauseOrderKindAttr::get(parser.getContext(), *enumValue);
902 return success();
903 }
904 return parser.emitError(loc, "invalid clause value: '") << enumStr << "'";
905}
906
908 ClauseOrderKindAttr order,
909 OrderModifierAttr orderMod) {
910 if (orderMod)
911 p << stringifyOrderModifier(orderMod.getValue()) << ":";
912 if (order)
913 p << stringifyClauseOrderKind(order.getValue());
914}
915
916template <typename ClauseTypeAttr, typename ClauseType>
917static ParseResult
918parseGranularityClause(OpAsmParser &parser, ClauseTypeAttr &prescriptiveness,
919 std::optional<OpAsmParser::UnresolvedOperand> &operand,
920 Type &operandType,
921 std::optional<ClauseType> (*symbolizeClause)(StringRef),
922 StringRef clauseName) {
923 StringRef enumStr;
924 if (succeeded(parser.parseOptionalKeyword(&enumStr))) {
925 if (std::optional<ClauseType> enumValue = symbolizeClause(enumStr)) {
926 prescriptiveness = ClauseTypeAttr::get(parser.getContext(), *enumValue);
927 if (parser.parseComma())
928 return failure();
929 } else {
930 return parser.emitError(parser.getCurrentLocation())
931 << "invalid " << clauseName << " modifier : '" << enumStr << "'";
932 ;
933 }
934 }
935
937 if (succeeded(parser.parseOperand(var))) {
938 operand = var;
939 } else {
940 return parser.emitError(parser.getCurrentLocation())
941 << "expected " << clauseName << " operand";
942 }
943
944 if (operand.has_value()) {
945 if (parser.parseColonType(operandType))
946 return failure();
947 }
948
949 return success();
950}
951
952template <typename ClauseTypeAttr, typename ClauseType>
953static void
955 ClauseTypeAttr prescriptiveness, Value operand,
956 mlir::Type operandType,
957 StringRef (*stringifyClauseType)(ClauseType)) {
958
959 if (prescriptiveness)
960 p << stringifyClauseType(prescriptiveness.getValue()) << ", ";
961
962 if (operand)
963 p << operand << ": " << operandType;
964}
965
966//===----------------------------------------------------------------------===//
967// Parser and printer for grainsize Clause
968//===----------------------------------------------------------------------===//
969
970// grainsize ::= `grainsize` `(` [strict ':'] grain-size `)`
971static ParseResult
972parseGrainsizeClause(OpAsmParser &parser, ClauseGrainsizeTypeAttr &grainsizeMod,
973 std::optional<OpAsmParser::UnresolvedOperand> &grainsize,
974 Type &grainsizeType) {
976 parser, grainsizeMod, grainsize, grainsizeType,
977 &symbolizeClauseGrainsizeType, "grainsize");
978}
979
981 ClauseGrainsizeTypeAttr grainsizeMod,
982 Value grainsize, mlir::Type grainsizeType) {
984 p, op, grainsizeMod, grainsize, grainsizeType,
985 &stringifyClauseGrainsizeType);
986}
987
988//===----------------------------------------------------------------------===//
989// Parser and printer for num_tasks Clause
990//===----------------------------------------------------------------------===//
991
992// numtask ::= `num_tasks` `(` [strict ':'] num-tasks `)`
993static ParseResult
994parseNumTasksClause(OpAsmParser &parser, ClauseNumTasksTypeAttr &numTasksMod,
995 std::optional<OpAsmParser::UnresolvedOperand> &numTasks,
996 Type &numTasksType) {
998 parser, numTasksMod, numTasks, numTasksType, &symbolizeClauseNumTasksType,
999 "num_tasks");
1000}
1001
1003 ClauseNumTasksTypeAttr numTasksMod,
1004 Value numTasks, mlir::Type numTasksType) {
1006 p, op, numTasksMod, numTasks, numTasksType, &stringifyClauseNumTasksType);
1007}
1008
1009//===----------------------------------------------------------------------===//
1010// Parser and printer for Heap Alloc Clause
1011//===----------------------------------------------------------------------===//
1012
1013/// operation ::= $in_type ( `(` $typeparams `)` )? ( `,` $shape )?
1014static ParseResult parseHeapAllocClause(
1015 OpAsmParser &parser, TypeAttr &inTypeAttr,
1017 SmallVectorImpl<Type> &typeparamsTypes,
1019 SmallVectorImpl<Type> &shapeTypes) {
1020 mlir::Type inType;
1021 if (parser.parseType(inType))
1022 return mlir::failure();
1023 inTypeAttr = TypeAttr::get(inType);
1024
1025 if (!parser.parseOptionalLParen()) {
1026 // parse the LEN params of the derived type. (<params> : <types>)
1027 if (parser.parseOperandList(typeparams, OpAsmParser::Delimiter::None) ||
1028 parser.parseColonTypeList(typeparamsTypes) || parser.parseRParen())
1029 return failure();
1030 }
1031
1032 if (!parser.parseOptionalComma()) {
1033 // parse size to scale by, vector of n dimensions of type index
1035 return failure();
1036
1037 // TODO: This overrides the actual types of the operands, which might cause
1038 // issues when they don't match. At the moment this is done in place of
1039 // making the corresponding operand type `Variadic<Index>` because index
1040 // types are lowered to I64 prior to LLVM IR translation.
1041 shapeTypes.append(shape.size(), IndexType::get(parser.getContext()));
1042 }
1043
1044 return success();
1045}
1046
1048 TypeAttr inType, ValueRange typeparams,
1049 TypeRange typeparamsTypes, ValueRange shape,
1050 TypeRange shapeTypes) {
1051 p << inType;
1052 if (!typeparams.empty()) {
1053 p << '(' << typeparams << " : " << typeparamsTypes << ')';
1054 }
1055 for (auto sh : shape) {
1056 p << ", ";
1057 p.printOperand(sh);
1058 }
1059}
1060
1061//===----------------------------------------------------------------------===//
1062// Parser, printer and verify for dyn_groupprivate Clause
1063//===----------------------------------------------------------------------===//
1064
1065static LogicalResult
1066verifyDynGroupprivateClause(Operation *op, AccessGroupModifierAttr accessGroup,
1067 FallbackModifierAttr fallback,
1068 Value dynGroupprivateSize) {
1069 if (!dynGroupprivateSize && (accessGroup || fallback))
1070 return op->emitOpError("dyn_groupprivate modifiers require a size operand");
1071
1072 return success();
1073}
1074
1076 OpAsmParser &parser, AccessGroupModifierAttr &accessGroupAttr,
1077 FallbackModifierAttr &fallbackAttr,
1078 std::optional<OpAsmParser::UnresolvedOperand> &dynGroupprivateSize,
1079 Type &sizeType) {
1080
1081 bool parsedAccessGroup = false;
1082 bool parsedFallback = false;
1083 bool parsedSize = false;
1084
1085 return parser.parseCommaSeparatedList([&]() -> ParseResult {
1086 // Parse AccessGroupModifier.
1087 if (succeeded(parser.parseOptionalKeyword("cgroup"))) {
1088 if (parsedAccessGroup)
1089 return parser.emitError(parser.getCurrentLocation(),
1090 "duplicate access group modifier");
1091 accessGroupAttr = AccessGroupModifierAttr::get(
1092 parser.getContext(), AccessGroupModifier::cgroup);
1093 parsedAccessGroup = true;
1094 return success();
1095 }
1096 // Parse FallbackModifier.
1097 if (succeeded(parser.parseOptionalKeyword("fallback"))) {
1098 if (parsedFallback)
1099 return parser.emitError(parser.getCurrentLocation(),
1100 "duplicate fallback modifier");
1101 if (parser.parseLParen())
1102 return parser.emitError(parser.getCurrentLocation(),
1103 "expected '(' after 'fallback'");
1104 llvm::StringRef fbKind;
1105 if (parser.parseKeyword(&fbKind))
1106 return parser.emitError(
1107 parser.getCurrentLocation(),
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;
1116 else
1117 return parser.emitError(parser.getCurrentLocation(),
1118 "invalid fallback modifier '" + fbKind + "'");
1119 fallbackAttr = FallbackModifierAttr::get(parser.getContext(), *fbEnum);
1120 if (parser.parseRParen())
1121 return parser.emitError(parser.getCurrentLocation(),
1122 "expected ')' after fallback modifier");
1123 parsedFallback = true;
1124 return success();
1125 }
1126 // Parse size operand.
1128 if (succeeded(parser.parseOperand(operand))) {
1129 if (parsedSize)
1130 return parser.emitError(parser.getCurrentLocation(),
1131 "duplicate size operand");
1132 dynGroupprivateSize = operand;
1133 parsedSize = true;
1134 if (failed(parser.parseColon()) || failed(parser.parseType(sizeType)))
1135 return parser.emitError(parser.getCurrentLocation(),
1136 "expected ':' and type after size operand");
1137 return success();
1138 }
1139 return parser.emitError(parser.getCurrentLocation(),
1140 "expected dyn_groupprivate_size operand");
1141 });
1142}
1143
1145 AccessGroupModifierAttr modifierFirst,
1146 FallbackModifierAttr modifierSecond,
1147 Value dynGroupprivateSize,
1148 Type sizeType) {
1149
1150 bool needsComma = false;
1151
1152 if (modifierFirst) {
1153 printer << modifierFirst.getValue();
1154 needsComma = true;
1155 }
1156
1157 if (modifierSecond) {
1158 if (needsComma)
1159 printer << ", ";
1160 printer << "fallback(";
1161 printer << modifierSecond.getValue();
1162 printer << ")";
1163 needsComma = true;
1164 }
1165
1166 if (dynGroupprivateSize) {
1167 if (needsComma)
1168 printer << ", ";
1169 printer << dynGroupprivateSize << " : " << sizeType;
1170 }
1171}
1172
1173//===----------------------------------------------------------------------===//
1174// Parser and printer for in_reduction Clause
1175//===----------------------------------------------------------------------===//
1176
1177/// Parses an `in_reduction` clause for an operation that does not give its
1178/// list items entry block arguments (e.g. `omp.target`). The expected format is
1179/// a comma-separated list of `[byref] @sym %var` followed by `: types`.
1180static ParseResult parseInReductionClause(
1181 OpAsmParser &parser,
1183 SmallVectorImpl<Type> &inReductionTypes,
1184 DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms) {
1186 SmallVector<bool> isByRefVec;
1187
1188 if (parser.parseCommaSeparatedList([&]() {
1189 isByRefVec.push_back(parser.parseOptionalKeyword("byref").succeeded());
1190 if (parser.parseAttribute(symbolVec.emplace_back()) ||
1191 parser.parseOperand(inReductionVars.emplace_back()))
1192 return failure();
1193 return success();
1194 }))
1195 return failure();
1196
1197 if (parser.parseColon())
1198 return failure();
1199
1200 if (parser.parseCommaSeparatedList(
1201 [&]() { return parser.parseType(inReductionTypes.emplace_back()); }))
1202 return failure();
1203
1204 if (inReductionVars.size() != inReductionTypes.size())
1205 return failure();
1206
1207 inReductionByref = makeDenseBoolArrayAttr(parser.getContext(), isByRefVec);
1208 SmallVector<Attribute> symbolAttrs(symbolVec.begin(), symbolVec.end());
1209 inReductionSyms = ArrayAttr::get(parser.getContext(), symbolAttrs);
1210 return success();
1211}
1212
1213/// Prints an `in_reduction` clause for an operation that does not give its list
1214/// items entry block arguments (e.g. `omp.target`). Mirrors
1215/// `parseInReductionClause`.
1217 ValueRange inReductionVars,
1218 TypeRange inReductionTypes,
1219 DenseBoolArrayAttr inReductionByref,
1220 ArrayAttr inReductionSyms) {
1221 MLIRContext *ctx = op->getContext();
1222
1223 ArrayAttr syms = inReductionSyms;
1224 if (!syms) {
1225 SmallVector<Attribute> values(inReductionVars.size(), nullptr);
1226 syms = ArrayAttr::get(ctx, values);
1227 }
1228
1229 DenseBoolArrayAttr byref = inReductionByref;
1230 if (!byref) {
1231 SmallVector<bool> values(inReductionVars.size(), false);
1232 byref = DenseBoolArrayAttr::get(ctx, values);
1233 }
1234
1235 llvm::interleaveComma(
1236 llvm::zip_equal(inReductionVars, syms.getValue(), byref.asArrayRef()), p,
1237 [&p](auto t) {
1238 auto [var, sym, isByRef] = t;
1239 if (isByRef)
1240 p << "byref ";
1241 if (sym)
1242 p << sym << " ";
1243 p << var;
1244 });
1245 p << " : ";
1246 llvm::interleaveComma(inReductionTypes, p);
1247}
1248
1249//===----------------------------------------------------------------------===//
1250// Parsers for operations including clauses that define entry block arguments.
1251//===----------------------------------------------------------------------===//
1252
1253namespace {
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) {}
1260};
1261struct PrivateParseArgs {
1262 llvm::SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars;
1263 llvm::SmallVectorImpl<Type> &types;
1264 ArrayAttr &syms;
1265 UnitAttr &needsBarrier;
1266 DenseI64ArrayAttr *mapIndices;
1267 PrivateParseArgs(SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars,
1268 SmallVectorImpl<Type> &types, ArrayAttr &syms,
1269 UnitAttr &needsBarrier,
1270 DenseI64ArrayAttr *mapIndices = nullptr)
1271 : vars(vars), types(types), syms(syms), needsBarrier(needsBarrier),
1272 mapIndices(mapIndices) {}
1273};
1274
1275struct ReductionParseArgs {
1276 SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars;
1277 SmallVectorImpl<Type> &types;
1278 DenseBoolArrayAttr &byref;
1279 ArrayAttr &syms;
1280 ReductionModifierAttr *modifier;
1281 ReductionParseArgs(SmallVectorImpl<OpAsmParser::UnresolvedOperand> &vars,
1282 SmallVectorImpl<Type> &types, DenseBoolArrayAttr &byref,
1283 ArrayAttr &syms, ReductionModifierAttr *mod = nullptr)
1284 : vars(vars), types(types), byref(byref), syms(syms), modifier(mod) {}
1285};
1286
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;
1297};
1298} // namespace
1299
1300static inline constexpr StringRef getPrivateNeedsBarrierSpelling() {
1301 return "private_barrier";
1302}
1303
1304static ParseResult parseClauseWithRegionArgs(
1305 OpAsmParser &parser,
1307 SmallVectorImpl<Type> &types,
1308 SmallVectorImpl<OpAsmParser::Argument> &regionPrivateArgs,
1309 ArrayAttr *symbols = nullptr, DenseI64ArrayAttr *mapIndices = nullptr,
1310 DenseBoolArrayAttr *byref = nullptr,
1311 ReductionModifierAttr *modifier = nullptr,
1312 UnitAttr *needsBarrier = nullptr) {
1314 SmallVector<int64_t> mapIndicesVec;
1315 SmallVector<bool> isByRefVec;
1316 unsigned regionArgOffset = regionPrivateArgs.size();
1317
1318 if (parser.parseLParen())
1319 return failure();
1320
1321 if (modifier && succeeded(parser.parseOptionalKeyword("mod"))) {
1322 StringRef enumStr;
1323 if (parser.parseColon() || parser.parseKeyword(&enumStr) ||
1324 parser.parseComma())
1325 return failure();
1326 std::optional<ReductionModifier> enumValue =
1327 symbolizeReductionModifier(enumStr);
1328 if (!enumValue.has_value())
1329 return failure();
1330 *modifier = ReductionModifierAttr::get(parser.getContext(), *enumValue);
1331 if (!*modifier)
1332 return failure();
1333 }
1334
1335 if (parser.parseCommaSeparatedList([&]() {
1336 if (byref)
1337 isByRefVec.push_back(
1338 parser.parseOptionalKeyword("byref").succeeded());
1339
1340 if (symbols && parser.parseAttribute(symbolVec.emplace_back()))
1341 return failure();
1342
1343 if (parser.parseOperand(operands.emplace_back()) ||
1344 parser.parseArrow() ||
1345 parser.parseArgument(regionPrivateArgs.emplace_back()))
1346 return failure();
1347
1348 if (mapIndices) {
1349 if (parser.parseOptionalLSquare().succeeded()) {
1350 if (parser.parseKeyword("map_idx") || parser.parseEqual() ||
1351 parser.parseInteger(mapIndicesVec.emplace_back()) ||
1352 parser.parseRSquare())
1353 return failure();
1354 } else {
1355 mapIndicesVec.push_back(-1);
1356 }
1357 }
1358
1359 return success();
1360 }))
1361 return failure();
1362
1363 if (parser.parseColon())
1364 return failure();
1365
1366 if (parser.parseCommaSeparatedList([&]() {
1367 if (parser.parseType(types.emplace_back()))
1368 return failure();
1369
1370 return success();
1371 }))
1372 return failure();
1373
1374 if (operands.size() != types.size())
1375 return failure();
1376
1377 if (parser.parseRParen())
1378 return failure();
1379
1380 if (needsBarrier) {
1382 .succeeded())
1383 *needsBarrier = mlir::UnitAttr::get(parser.getContext());
1384 }
1385
1386 auto *argsBegin = regionPrivateArgs.begin();
1387 MutableArrayRef argsSubrange(argsBegin + regionArgOffset,
1388 argsBegin + regionArgOffset + types.size());
1389 for (auto [prv, type] : llvm::zip_equal(argsSubrange, types)) {
1390 prv.type = type;
1391 }
1392
1393 if (symbols) {
1394 SmallVector<Attribute> symbolAttrs(symbolVec.begin(), symbolVec.end());
1395 *symbols = ArrayAttr::get(parser.getContext(), symbolAttrs);
1396 }
1397
1398 if (!mapIndicesVec.empty())
1399 *mapIndices =
1400 mlir::DenseI64ArrayAttr::get(parser.getContext(), mapIndicesVec);
1401
1402 if (byref)
1403 *byref = makeDenseBoolArrayAttr(parser.getContext(), isByRefVec);
1404
1405 return success();
1406}
1407
1408static ParseResult parseBlockArgClause(
1409 OpAsmParser &parser,
1411 StringRef keyword, std::optional<MapParseArgs> mapArgs) {
1412 if (succeeded(parser.parseOptionalKeyword(keyword))) {
1413 if (!mapArgs)
1414 return failure();
1415
1416 if (failed(parseClauseWithRegionArgs(parser, mapArgs->vars, mapArgs->types,
1417 entryBlockArgs)))
1418 return failure();
1419 }
1420 return success();
1421}
1422
1423static ParseResult parseBlockArgClause(
1424 OpAsmParser &parser,
1426 StringRef keyword, std::optional<PrivateParseArgs> privateArgs) {
1427 if (succeeded(parser.parseOptionalKeyword(keyword))) {
1428 if (!privateArgs)
1429 return failure();
1430
1431 if (failed(parseClauseWithRegionArgs(
1432 parser, privateArgs->vars, privateArgs->types, entryBlockArgs,
1433 &privateArgs->syms, privateArgs->mapIndices, /*byref=*/nullptr,
1434 /*modifier=*/nullptr, &privateArgs->needsBarrier)))
1435 return failure();
1436 }
1437 return success();
1438}
1439
1440static ParseResult parseBlockArgClause(
1441 OpAsmParser &parser,
1443 StringRef keyword, std::optional<ReductionParseArgs> reductionArgs) {
1444 if (succeeded(parser.parseOptionalKeyword(keyword))) {
1445 if (!reductionArgs)
1446 return failure();
1447 if (failed(parseClauseWithRegionArgs(
1448 parser, reductionArgs->vars, reductionArgs->types, entryBlockArgs,
1449 &reductionArgs->syms, /*mapIndices=*/nullptr, &reductionArgs->byref,
1450 reductionArgs->modifier)))
1451 return failure();
1452 }
1453 return success();
1454}
1455
1456static ParseResult parseBlockArgRegion(OpAsmParser &parser, Region &region,
1457 AllRegionParseArgs args) {
1459
1460 if (failed(parseBlockArgClause(parser, entryBlockArgs, "has_device_addr",
1461 args.hasDeviceAddrArgs)))
1462 return parser.emitError(parser.getCurrentLocation())
1463 << "invalid `has_device_addr` format";
1464
1465 if (failed(parseBlockArgClause(parser, entryBlockArgs, "host_eval",
1466 args.hostEvalArgs)))
1467 return parser.emitError(parser.getCurrentLocation())
1468 << "invalid `host_eval` format";
1469
1470 if (failed(parseBlockArgClause(parser, entryBlockArgs, "in_reduction",
1471 args.inReductionArgs)))
1472 return parser.emitError(parser.getCurrentLocation())
1473 << "invalid `in_reduction` format";
1474
1475 if (failed(parseBlockArgClause(parser, entryBlockArgs, "map_entries",
1476 args.mapArgs)))
1477 return parser.emitError(parser.getCurrentLocation())
1478 << "invalid `map_entries` format";
1479
1480 if (failed(parseBlockArgClause(parser, entryBlockArgs, "private",
1481 args.privateArgs)))
1482 return parser.emitError(parser.getCurrentLocation())
1483 << "invalid `private` format";
1484
1485 if (failed(parseBlockArgClause(parser, entryBlockArgs, "reduction",
1486 args.reductionArgs)))
1487 return parser.emitError(parser.getCurrentLocation())
1488 << "invalid `reduction` format";
1489
1490 if (failed(parseBlockArgClause(parser, entryBlockArgs, "task_reduction",
1491 args.taskReductionArgs)))
1492 return parser.emitError(parser.getCurrentLocation())
1493 << "invalid `task_reduction` format";
1494
1495 if (failed(parseBlockArgClause(parser, entryBlockArgs, "use_device_addr",
1496 args.useDeviceAddrArgs)))
1497 return parser.emitError(parser.getCurrentLocation())
1498 << "invalid `use_device_addr` format";
1499
1500 if (failed(parseBlockArgClause(parser, entryBlockArgs, "use_device_ptr",
1501 args.useDevicePtrArgs)))
1502 return parser.emitError(parser.getCurrentLocation())
1503 << "invalid `use_device_addr` format";
1504
1505 return parser.parseRegion(region, entryBlockArgs);
1506}
1507
1508// These parseXyz functions correspond to the custom<Xyz> definitions
1509// in the .td file(s).
1510static ParseResult parseTargetOpRegion(
1511 OpAsmParser &parser, Region &region,
1513 SmallVectorImpl<Type> &hasDeviceAddrTypes,
1515 SmallVectorImpl<Type> &hostEvalTypes,
1517 SmallVectorImpl<Type> &mapTypes,
1519 llvm::SmallVectorImpl<Type> &privateTypes, ArrayAttr &privateSyms,
1520 UnitAttr &privateNeedsBarrier, DenseI64ArrayAttr &privateMaps) {
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);
1527 return parseBlockArgRegion(parser, region, args);
1528}
1529
1531 OpAsmParser &parser, Region &region,
1533 SmallVectorImpl<Type> &inReductionTypes,
1534 DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms,
1536 llvm::SmallVectorImpl<Type> &privateTypes, ArrayAttr &privateSyms,
1537 UnitAttr &privateNeedsBarrier) {
1538 AllRegionParseArgs args;
1539 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1540 inReductionByref, inReductionSyms);
1541 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1542 privateNeedsBarrier);
1543 return parseBlockArgRegion(parser, region, args);
1544}
1545
1547 OpAsmParser &parser, Region &region,
1549 SmallVectorImpl<Type> &inReductionTypes,
1550 DenseBoolArrayAttr &inReductionByref, ArrayAttr &inReductionSyms,
1552 llvm::SmallVectorImpl<Type> &privateTypes, ArrayAttr &privateSyms,
1553 UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod,
1555 SmallVectorImpl<Type> &reductionTypes, DenseBoolArrayAttr &reductionByref,
1556 ArrayAttr &reductionSyms) {
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);
1564 return parseBlockArgRegion(parser, region, args);
1565}
1566
1567static ParseResult parsePrivateRegion(
1568 OpAsmParser &parser, Region &region,
1570 llvm::SmallVectorImpl<Type> &privateTypes, ArrayAttr &privateSyms,
1571 UnitAttr &privateNeedsBarrier) {
1572 AllRegionParseArgs args;
1573 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1574 privateNeedsBarrier);
1575 return parseBlockArgRegion(parser, region, args);
1576}
1577
1579 OpAsmParser &parser, Region &region,
1581 llvm::SmallVectorImpl<Type> &privateTypes, ArrayAttr &privateSyms,
1582 UnitAttr &privateNeedsBarrier, ReductionModifierAttr &reductionMod,
1584 SmallVectorImpl<Type> &reductionTypes, DenseBoolArrayAttr &reductionByref,
1585 ArrayAttr &reductionSyms) {
1586 AllRegionParseArgs args;
1587 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1588 privateNeedsBarrier);
1589 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1590 reductionSyms, &reductionMod);
1591 return parseBlockArgRegion(parser, region, args);
1592}
1593
1594static ParseResult parseTaskReductionRegion(
1595 OpAsmParser &parser, Region &region,
1597 SmallVectorImpl<Type> &taskReductionTypes,
1598 DenseBoolArrayAttr &taskReductionByref, ArrayAttr &taskReductionSyms) {
1599 AllRegionParseArgs args;
1600 args.taskReductionArgs.emplace(taskReductionVars, taskReductionTypes,
1601 taskReductionByref, taskReductionSyms);
1602 return parseBlockArgRegion(parser, region, args);
1603}
1604
1606 OpAsmParser &parser, Region &region,
1608 SmallVectorImpl<Type> &useDeviceAddrTypes,
1610 SmallVectorImpl<Type> &useDevicePtrTypes) {
1611 AllRegionParseArgs args;
1612 args.useDeviceAddrArgs.emplace(useDeviceAddrVars, useDeviceAddrTypes);
1613 args.useDevicePtrArgs.emplace(useDevicePtrVars, useDevicePtrTypes);
1614 return parseBlockArgRegion(parser, region, args);
1615}
1616
1617//===----------------------------------------------------------------------===//
1618// Printers for operations including clauses that define entry block arguments.
1619//===----------------------------------------------------------------------===//
1620
1621namespace {
1622struct MapPrintArgs {
1623 ValueRange vars;
1624 TypeRange types;
1625 MapPrintArgs(ValueRange vars, TypeRange types) : vars(vars), types(types) {}
1626};
1627struct PrivatePrintArgs {
1628 ValueRange vars;
1629 TypeRange types;
1630 ArrayAttr syms;
1631 UnitAttr needsBarrier;
1632 DenseI64ArrayAttr mapIndices;
1633 PrivatePrintArgs(ValueRange vars, TypeRange types, ArrayAttr syms,
1634 UnitAttr needsBarrier, DenseI64ArrayAttr mapIndices)
1635 : vars(vars), types(types), syms(syms), needsBarrier(needsBarrier),
1636 mapIndices(mapIndices) {}
1637};
1638struct ReductionPrintArgs {
1639 ValueRange vars;
1640 TypeRange types;
1641 DenseBoolArrayAttr byref;
1642 ArrayAttr syms;
1643 ReductionModifierAttr modifier;
1644 ReductionPrintArgs(ValueRange vars, TypeRange types, DenseBoolArrayAttr byref,
1645 ArrayAttr syms, ReductionModifierAttr mod = nullptr)
1646 : vars(vars), types(types), byref(byref), syms(syms), modifier(mod) {}
1647};
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;
1658};
1659} // namespace
1660
1662 OpAsmPrinter &p, MLIRContext *ctx, StringRef clauseName,
1663 ValueRange argsSubrange, ValueRange operands, TypeRange types,
1664 ArrayAttr symbols = nullptr, DenseI64ArrayAttr mapIndices = nullptr,
1665 DenseBoolArrayAttr byref = nullptr,
1666 ReductionModifierAttr modifier = nullptr, UnitAttr needsBarrier = nullptr) {
1667 if (argsSubrange.empty())
1668 return;
1669
1670 p << clauseName << "(";
1671
1672 if (modifier)
1673 p << "mod: " << stringifyReductionModifier(modifier.getValue()) << ", ";
1674
1675 if (!symbols) {
1676 llvm::SmallVector<Attribute> values(operands.size(), nullptr);
1677 symbols = ArrayAttr::get(ctx, values);
1678 }
1679
1680 if (!mapIndices) {
1681 llvm::SmallVector<int64_t> values(operands.size(), -1);
1682 mapIndices = DenseI64ArrayAttr::get(ctx, values);
1683 }
1684
1685 if (!byref) {
1686 mlir::SmallVector<bool> values(operands.size(), false);
1687 byref = DenseBoolArrayAttr::get(ctx, values);
1688 }
1689
1690 llvm::interleaveComma(llvm::zip_equal(operands, argsSubrange, symbols,
1691 mapIndices.asArrayRef(),
1692 byref.asArrayRef()),
1693 p, [&p](auto t) {
1694 auto [op, arg, sym, map, isByRef] = t;
1695 if (isByRef)
1696 p << "byref ";
1697 if (sym)
1698 p << sym << " ";
1699
1700 p << op << " -> " << arg;
1701
1702 if (map != -1)
1703 p << " [map_idx=" << map << "]";
1704 });
1705 p << " : ";
1706 llvm::interleaveComma(types, p);
1707 p << ") ";
1708
1709 if (needsBarrier)
1710 p << getPrivateNeedsBarrierSpelling() << " ";
1711}
1712
1714 StringRef clauseName, ValueRange argsSubrange,
1715 std::optional<MapPrintArgs> mapArgs) {
1716 if (mapArgs)
1717 printClauseWithRegionArgs(p, ctx, clauseName, argsSubrange, mapArgs->vars,
1718 mapArgs->types);
1719}
1720
1722 StringRef clauseName, ValueRange argsSubrange,
1723 std::optional<PrivatePrintArgs> privateArgs) {
1724 if (privateArgs)
1726 p, ctx, clauseName, argsSubrange, privateArgs->vars, privateArgs->types,
1727 privateArgs->syms, privateArgs->mapIndices, /*byref=*/nullptr,
1728 /*modifier=*/nullptr, privateArgs->needsBarrier);
1729}
1730
1731static void
1732printBlockArgClause(OpAsmPrinter &p, MLIRContext *ctx, StringRef clauseName,
1733 ValueRange argsSubrange,
1734 std::optional<ReductionPrintArgs> reductionArgs) {
1735 if (reductionArgs)
1736 printClauseWithRegionArgs(p, ctx, clauseName, argsSubrange,
1737 reductionArgs->vars, reductionArgs->types,
1738 reductionArgs->syms, /*mapIndices=*/nullptr,
1739 reductionArgs->byref, reductionArgs->modifier);
1740}
1741
1743 const AllRegionPrintArgs &args) {
1744 auto iface = llvm::cast<mlir::omp::BlockArgOpenMPOpInterface>(op);
1745 MLIRContext *ctx = op->getContext();
1746
1747 printBlockArgClause(p, ctx, "has_device_addr",
1748 iface.getHasDeviceAddrBlockArgs(),
1749 args.hasDeviceAddrArgs);
1750 printBlockArgClause(p, ctx, "host_eval", iface.getHostEvalBlockArgs(),
1751 args.hostEvalArgs);
1752 printBlockArgClause(p, ctx, "in_reduction", iface.getInReductionBlockArgs(),
1753 args.inReductionArgs);
1754 printBlockArgClause(p, ctx, "map_entries", iface.getMapBlockArgs(),
1755 args.mapArgs);
1756 printBlockArgClause(p, ctx, "private", iface.getPrivateBlockArgs(),
1757 args.privateArgs);
1758 printBlockArgClause(p, ctx, "reduction", iface.getReductionBlockArgs(),
1759 args.reductionArgs);
1760 printBlockArgClause(p, ctx, "task_reduction",
1761 iface.getTaskReductionBlockArgs(),
1762 args.taskReductionArgs);
1763 printBlockArgClause(p, ctx, "use_device_addr",
1764 iface.getUseDeviceAddrBlockArgs(),
1765 args.useDeviceAddrArgs);
1766 printBlockArgClause(p, ctx, "use_device_ptr",
1767 iface.getUseDevicePtrBlockArgs(), args.useDevicePtrArgs);
1768
1769 p.printRegion(region, /*printEntryBlockArgs=*/false);
1770}
1771
1772// These parseXyz functions correspond to the custom<Xyz> definitions
1773// in the .td file(s).
1775 ValueRange hasDeviceAddrVars,
1776 TypeRange hasDeviceAddrTypes,
1777 ValueRange hostEvalVars,
1778 TypeRange hostEvalTypes, ValueRange mapVars,
1779 TypeRange mapTypes, ValueRange privateVars,
1780 TypeRange privateTypes, ArrayAttr privateSyms,
1781 UnitAttr privateNeedsBarrier,
1782 DenseI64ArrayAttr privateMaps) {
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);
1789 printBlockArgRegion(p, op, region, args);
1790}
1791
1793 OpAsmPrinter &p, Operation *op, Region &region, ValueRange inReductionVars,
1794 TypeRange inReductionTypes, DenseBoolArrayAttr inReductionByref,
1795 ArrayAttr inReductionSyms, ValueRange privateVars, TypeRange privateTypes,
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,
1802 /*mapIndices=*/nullptr);
1803 printBlockArgRegion(p, op, region, args);
1804}
1805
1807 OpAsmPrinter &p, Operation *op, Region &region, ValueRange inReductionVars,
1808 TypeRange inReductionTypes, DenseBoolArrayAttr inReductionByref,
1809 ArrayAttr inReductionSyms, ValueRange privateVars, TypeRange privateTypes,
1810 ArrayAttr privateSyms, UnitAttr privateNeedsBarrier,
1811 ReductionModifierAttr reductionMod, ValueRange reductionVars,
1812 TypeRange reductionTypes, DenseBoolArrayAttr reductionByref,
1813 ArrayAttr reductionSyms) {
1814 AllRegionPrintArgs args;
1815 args.inReductionArgs.emplace(inReductionVars, inReductionTypes,
1816 inReductionByref, inReductionSyms);
1817 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1818 privateNeedsBarrier,
1819 /*mapIndices=*/nullptr);
1820 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1821 reductionSyms, reductionMod);
1822 printBlockArgRegion(p, op, region, args);
1823}
1824
1826 ValueRange privateVars, TypeRange privateTypes,
1827 ArrayAttr privateSyms,
1828 UnitAttr privateNeedsBarrier) {
1829 AllRegionPrintArgs args;
1830 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1831 privateNeedsBarrier,
1832 /*mapIndices=*/nullptr);
1833 printBlockArgRegion(p, op, region, args);
1834}
1835
1837 OpAsmPrinter &p, Operation *op, Region &region, ValueRange privateVars,
1838 TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier,
1839 ReductionModifierAttr reductionMod, ValueRange reductionVars,
1840 TypeRange reductionTypes, DenseBoolArrayAttr reductionByref,
1841 ArrayAttr reductionSyms) {
1842 AllRegionPrintArgs args;
1843 args.privateArgs.emplace(privateVars, privateTypes, privateSyms,
1844 privateNeedsBarrier,
1845 /*mapIndices=*/nullptr);
1846 args.reductionArgs.emplace(reductionVars, reductionTypes, reductionByref,
1847 reductionSyms, reductionMod);
1848 printBlockArgRegion(p, op, region, args);
1849}
1850
1852 Region &region,
1853 ValueRange taskReductionVars,
1854 TypeRange taskReductionTypes,
1855 DenseBoolArrayAttr taskReductionByref,
1856 ArrayAttr taskReductionSyms) {
1857 AllRegionPrintArgs args;
1858 args.taskReductionArgs.emplace(taskReductionVars, taskReductionTypes,
1859 taskReductionByref, taskReductionSyms);
1860 printBlockArgRegion(p, op, region, args);
1861}
1862
1864 Region &region,
1865 ValueRange useDeviceAddrVars,
1866 TypeRange useDeviceAddrTypes,
1867 ValueRange useDevicePtrVars,
1868 TypeRange useDevicePtrTypes) {
1869 AllRegionPrintArgs args;
1870 args.useDeviceAddrArgs.emplace(useDeviceAddrVars, useDeviceAddrTypes);
1871 args.useDevicePtrArgs.emplace(useDevicePtrVars, useDevicePtrTypes);
1872 printBlockArgRegion(p, op, region, args);
1873}
1874
1875template <typename ParsePrefixFn>
1876static ParseResult parseSplitIteratedList(
1877 OpAsmParser &parser,
1879 SmallVectorImpl<Type> &iteratedTypes,
1881 SmallVectorImpl<Type> &plainTypes, ParsePrefixFn &&parsePrefix) {
1882
1883 return parser.parseCommaSeparatedList([&]() -> ParseResult {
1884 if (failed(parsePrefix()))
1885 return failure();
1886
1888 Type ty;
1889 if (parser.parseOperand(v) || parser.parseColonType(ty))
1890 return failure();
1891
1892 if (llvm::isa<mlir::omp::IteratedType>(ty)) {
1893 iteratedVars.push_back(v);
1894 iteratedTypes.push_back(ty);
1895 } else {
1896 plainVars.push_back(v);
1897 plainTypes.push_back(ty);
1898 }
1899 return success();
1900 });
1901}
1902
1903template <typename PrintPrefixFn>
1905 TypeRange iteratedTypes,
1906 ValueRange plainVars, TypeRange plainTypes,
1907 PrintPrefixFn &&printPrefixForPlain,
1908 PrintPrefixFn &&printPrefixForIterated) {
1909
1910 bool first = true;
1911 auto emit = [&](Value v, Type t, auto &&printPrefix) {
1912 if (!first)
1913 p << ", ";
1914 printPrefix(v, t);
1915 p << v << " : " << t;
1916 first = false;
1917 };
1918
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);
1923}
1924
1925/// Verifies Reduction Clause
1926static LogicalResult
1927verifyReductionVarList(Operation *op, std::optional<ArrayAttr> reductionSyms,
1928 OperandRange reductionVars,
1929 std::optional<ArrayRef<bool>> reductionByref) {
1930 if (!reductionVars.empty()) {
1931 if (!reductionSyms || reductionSyms->size() != reductionVars.size())
1932 return op->emitOpError()
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";
1938 } else {
1939 if (reductionSyms)
1940 return op->emitOpError() << "unexpected reduction symbol references";
1941 return success();
1942 }
1943
1944 // TODO: The followings should be done in
1945 // SymbolUserOpInterface::verifySymbolUses.
1946 DenseSet<Value> accumulators;
1947 for (auto args : llvm::zip(reductionVars, *reductionSyms)) {
1948 Value accum = std::get<0>(args);
1949
1950 if (!accumulators.insert(accum).second)
1951 return op->emitOpError() << "accumulator variable used more than once";
1952
1953 Type varType = accum.getType();
1954 auto symbolRef = llvm::cast<SymbolRefAttr>(std::get<1>(args));
1955 auto decl =
1957 if (!decl)
1958 return op->emitOpError() << "expected symbol reference " << symbolRef
1959 << " to point to a reduction declaration";
1960
1961 if (decl.getAccumulatorType() && decl.getAccumulatorType() != varType)
1962 return op->emitOpError()
1963 << "expected accumulator (" << varType
1964 << ") to be the same type as reduction declaration ("
1965 << decl.getAccumulatorType() << ")";
1966 }
1967
1968 return success();
1969}
1970
1971//===----------------------------------------------------------------------===//
1972// Parser, printer and verifier for Copyprivate
1973//===----------------------------------------------------------------------===//
1974
1975/// copyprivate-entry-list ::= copyprivate-entry
1976/// | copyprivate-entry-list `,` copyprivate-entry
1977/// copyprivate-entry ::= ssa-id `->` symbol-ref `:` type
1978static ParseResult parseCopyprivate(
1979 OpAsmParser &parser,
1981 SmallVectorImpl<Type> &copyprivateTypes, ArrayAttr &copyprivateSyms) {
1983 if (failed(parser.parseCommaSeparatedList([&]() {
1984 if (parser.parseOperand(copyprivateVars.emplace_back()) ||
1985 parser.parseArrow() ||
1986 parser.parseAttribute(symsVec.emplace_back()) ||
1987 parser.parseColonType(copyprivateTypes.emplace_back()))
1988 return failure();
1989 return success();
1990 })))
1991 return failure();
1992 SmallVector<Attribute> syms(symsVec.begin(), symsVec.end());
1993 copyprivateSyms = ArrayAttr::get(parser.getContext(), syms);
1994 return success();
1995}
1996
1997/// Print Copyprivate clause
1999 OperandRange copyprivateVars,
2000 TypeRange copyprivateTypes,
2001 std::optional<ArrayAttr> copyprivateSyms) {
2002 if (!copyprivateSyms.has_value())
2003 return;
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);
2009 });
2010}
2011
2012/// Verifies CopyPrivate Clause
2013static LogicalResult
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())
2024 return success();
2025
2026 for (auto copyprivateVarAndSym :
2027 llvm::zip(copyprivateVars, *copyprivateSyms)) {
2028 auto symbolRef =
2029 llvm::cast<SymbolRefAttr>(std::get<1>(copyprivateVarAndSym));
2030 std::optional<std::variant<mlir::func::FuncOp, mlir::LLVM::LLVMFuncOp>>
2031 funcOp;
2032 if (mlir::func::FuncOp mlirFuncOp =
2034 symbolRef))
2035 funcOp = mlirFuncOp;
2036 else if (mlir::LLVM::LLVMFuncOp llvmFuncOp =
2038 op, symbolRef))
2039 funcOp = llvmFuncOp;
2040
2041 auto getNumArguments = [&] {
2042 return std::visit([](auto &f) { return f.getNumArguments(); }, *funcOp);
2043 };
2044
2045 auto getArgumentType = [&](unsigned i) {
2046 return std::visit([i](auto &f) { return f.getArgumentTypes()[i]; },
2047 *funcOp);
2048 };
2049
2050 if (!funcOp)
2051 return op->emitOpError() << "expected symbol reference " << symbolRef
2052 << " to point to a copy function";
2053
2054 if (getNumArguments() != 2)
2055 return op->emitOpError()
2056 << "expected copy function " << symbolRef << " to have 2 operands";
2057
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";
2062
2063 Type varType = std::get<0>(copyprivateVarAndSym).getType();
2064 if (argTy != varType)
2065 return op->emitOpError()
2066 << "expected copy function arguments' type (" << argTy
2067 << ") to be the same as copyprivate variable's type (" << varType
2068 << ")";
2069 }
2070
2071 return success();
2072}
2073
2074//===----------------------------------------------------------------------===//
2075// Parser, printer and verifier for DependVarList
2076//===----------------------------------------------------------------------===//
2077
2078/// depend-entry-list ::= depend-entry
2079/// | depend-entry-list `,` depend-entry
2080/// depend-entry ::= depend-kind `->` ssa-id `:` type
2081/// | depend-kind `->` ssa-id `:` iterated-type
2082static ParseResult parseDependVarList(
2083 OpAsmParser &parser,
2085 SmallVectorImpl<Type> &dependTypes, ArrayAttr &dependKinds,
2087 SmallVectorImpl<Type> &iteratedTypes, ArrayAttr &iteratedKinds) {
2090 if (failed(parser.parseCommaSeparatedList([&]() {
2091 StringRef keyword;
2092 OpAsmParser::UnresolvedOperand operand;
2093 Type ty;
2094 if (parser.parseKeyword(&keyword) || parser.parseArrow() ||
2095 parser.parseOperand(operand) || parser.parseColonType(ty))
2096 return failure();
2097 std::optional<ClauseTaskDepend> keywordDepend =
2098 symbolizeClauseTaskDepend(keyword);
2099 if (!keywordDepend)
2100 return failure();
2101 auto kindAttr =
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);
2107 } else {
2108 dependVars.push_back(operand);
2109 dependTypes.push_back(ty);
2110 kindsVec.push_back(kindAttr);
2111 }
2112 return success();
2113 })))
2114 return failure();
2115 SmallVector<Attribute> kinds(kindsVec.begin(), kindsVec.end());
2116 dependKinds = ArrayAttr::get(parser.getContext(), kinds);
2117 SmallVector<Attribute> iterKinds(iterKindsVec.begin(), iterKindsVec.end());
2118 iteratedKinds = ArrayAttr::get(parser.getContext(), iterKinds);
2119 return success();
2120}
2121
2122/// Print Depend clause
2124 OperandRange dependVars, TypeRange dependTypes,
2125 std::optional<ArrayAttr> dependKinds,
2126 OperandRange iteratedVars,
2127 TypeRange iteratedTypes,
2128 std::optional<ArrayAttr> iteratedKinds) {
2129 bool first = true;
2130 auto printEntries = [&](OperandRange vars, TypeRange types,
2131 std::optional<ArrayAttr> kinds) {
2132 for (unsigned i = 0, e = vars.size(); i < e; ++i) {
2133 if (!first)
2134 p << ", ";
2135 p << stringifyClauseTaskDepend(
2136 llvm::cast<mlir::omp::ClauseTaskDependAttr>((*kinds)[i])
2137 .getValue())
2138 << " -> " << vars[i] << " : " << types[i];
2139 first = false;
2140 }
2141 };
2142 printEntries(dependVars, dependTypes, dependKinds);
2143 printEntries(iteratedVars, iteratedTypes, iteratedKinds);
2144}
2145
2146/// Verifies Depend clause
2147static LogicalResult verifyDependVarList(Operation *op,
2148 std::optional<ArrayAttr> dependKinds,
2149 OperandRange dependVars,
2150 std::optional<ArrayAttr> iteratedKinds,
2151 OperandRange iteratedVars) {
2152 if (!dependVars.empty()) {
2153 if (!dependKinds || dependKinds->size() != dependVars.size())
2154 return op->emitOpError() << "expected as many depend values"
2155 " as depend variables";
2156 } else {
2157 if (dependKinds && !dependKinds->empty())
2158 return op->emitOpError() << "unexpected depend values";
2159 }
2160
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";
2165 } else {
2166 if (iteratedKinds && !iteratedKinds->empty())
2167 return op->emitOpError() << "unexpected depend iterated values";
2168 }
2169
2170 return success();
2171}
2172
2173//===----------------------------------------------------------------------===//
2174// Parser, printer and verifier for Synchronization Hint (2.17.12)
2175//===----------------------------------------------------------------------===//
2176
2177/// Parses a Synchronization Hint clause. The value of hint is an integer
2178/// which is a combination of different hints from `omp_sync_hint_t`.
2179///
2180/// hint-clause = `hint` `(` hint-value `)`
2181static ParseResult parseSynchronizationHint(OpAsmParser &parser,
2182 IntegerAttr &hintAttr) {
2183 StringRef hintKeyword;
2184 int64_t hint = 0;
2185 if (succeeded(parser.parseOptionalKeyword("none"))) {
2186 hintAttr = IntegerAttr::get(parser.getBuilder().getI64Type(), 0);
2187 return success();
2188 }
2189 auto parseKeyword = [&]() -> ParseResult {
2190 if (failed(parser.parseKeyword(&hintKeyword)))
2191 return failure();
2192 if (hintKeyword == "uncontended")
2193 hint |= 1;
2194 else if (hintKeyword == "contended")
2195 hint |= 2;
2196 else if (hintKeyword == "nonspeculative")
2197 hint |= 4;
2198 else if (hintKeyword == "speculative")
2199 hint |= 8;
2200 else
2201 return parser.emitError(parser.getCurrentLocation())
2202 << hintKeyword << " is not a valid hint";
2203 return success();
2204 };
2205 if (parser.parseCommaSeparatedList(parseKeyword))
2206 return failure();
2207 hintAttr = IntegerAttr::get(parser.getBuilder().getI64Type(), hint);
2208 return success();
2209}
2210
2211/// Prints a Synchronization Hint clause
2213 IntegerAttr hintAttr) {
2214 int64_t hint = hintAttr.getInt();
2215
2216 if (hint == 0) {
2217 p << "none";
2218 return;
2219 }
2220
2221 // Helper function to get n-th bit from the right end of `value`
2222 auto bitn = [](int value, int n) -> bool { return value & (1 << n); };
2223
2224 bool uncontended = bitn(hint, 0);
2225 bool contended = bitn(hint, 1);
2226 bool nonspeculative = bitn(hint, 2);
2227 bool speculative = bitn(hint, 3);
2228
2230 if (uncontended)
2231 hints.push_back("uncontended");
2232 if (contended)
2233 hints.push_back("contended");
2234 if (nonspeculative)
2235 hints.push_back("nonspeculative");
2236 if (speculative)
2237 hints.push_back("speculative");
2238
2239 llvm::interleaveComma(hints, p);
2240}
2241
2242/// Verifies a synchronization hint clause
2243static LogicalResult verifySynchronizationHint(Operation *op, uint64_t hint) {
2244
2245 // Helper function to get n-th bit from the right end of `value`
2246 auto bitn = [](int value, int n) -> bool { return value & (1 << n); };
2247
2248 bool uncontended = bitn(hint, 0);
2249 bool contended = bitn(hint, 1);
2250 bool nonspeculative = bitn(hint, 2);
2251 bool speculative = bitn(hint, 3);
2252
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.";
2259 return success();
2260}
2261
2262//===----------------------------------------------------------------------===//
2263// Parser, printer and verifier for Target
2264//===----------------------------------------------------------------------===//
2265
2266// Helper function to get bitwise AND of `value` and 'flag' then return it as a
2267// boolean
2268static bool mapTypeToBool(ClauseMapFlags value, ClauseMapFlags flag) {
2269 return (value & flag) == flag;
2270}
2271
2272/// Parses a map_entries map type from a string format back into its numeric
2273/// value.
2274///
2275/// map-clause = `map_clauses ( ( `(` `always, `? `implicit, `? `ompx_hold, `?
2276/// `close, `? `present, `? ( `to` | `from` | `delete` `)` )+ `)` )
2277static ParseResult parseMapClause(OpAsmParser &parser,
2278 ClauseMapFlagsAttr &mapType) {
2279 ClauseMapFlags mapTypeBits = ClauseMapFlags::none;
2280 // This simply verifies the correct keyword is read in, the
2281 // keyword itself is stored inside of the operation
2282 auto parseTypeAndMod = [&]() -> ParseResult {
2283 StringRef mapTypeMod;
2284 if (parser.parseKeyword(&mapTypeMod))
2285 return failure();
2286
2287 if (mapTypeMod == "always")
2288 mapTypeBits |= ClauseMapFlags::always;
2289
2290 if (mapTypeMod == "implicit")
2291 mapTypeBits |= ClauseMapFlags::implicit;
2292
2293 if (mapTypeMod == "ompx_hold")
2294 mapTypeBits |= ClauseMapFlags::ompx_hold;
2295
2296 if (mapTypeMod == "close")
2297 mapTypeBits |= ClauseMapFlags::close;
2298
2299 if (mapTypeMod == "present")
2300 mapTypeBits |= ClauseMapFlags::present;
2301
2302 if (mapTypeMod == "to")
2303 mapTypeBits |= ClauseMapFlags::to;
2304
2305 if (mapTypeMod == "from")
2306 mapTypeBits |= ClauseMapFlags::from;
2307
2308 if (mapTypeMod == "tofrom")
2309 mapTypeBits |= ClauseMapFlags::to | ClauseMapFlags::from;
2310
2311 if (mapTypeMod == "delete")
2312 mapTypeBits |= ClauseMapFlags::del;
2313
2314 if (mapTypeMod == "storage")
2315 mapTypeBits |= ClauseMapFlags::storage;
2316
2317 if (mapTypeMod == "return_param")
2318 mapTypeBits |= ClauseMapFlags::return_param;
2319
2320 if (mapTypeMod == "private")
2321 mapTypeBits |= ClauseMapFlags::priv;
2322
2323 if (mapTypeMod == "literal")
2324 mapTypeBits |= ClauseMapFlags::literal;
2325
2326 if (mapTypeMod == "attach")
2327 mapTypeBits |= ClauseMapFlags::attach;
2328
2329 if (mapTypeMod == "attach_always")
2330 mapTypeBits |= ClauseMapFlags::attach_always;
2331
2332 if (mapTypeMod == "attach_never")
2333 mapTypeBits |= ClauseMapFlags::attach_never;
2334
2335 if (mapTypeMod == "attach_auto")
2336 mapTypeBits |= ClauseMapFlags::attach_auto;
2337
2338 if (mapTypeMod == "ref_ptr")
2339 mapTypeBits |= ClauseMapFlags::ref_ptr;
2340
2341 if (mapTypeMod == "ref_ptee")
2342 mapTypeBits |= ClauseMapFlags::ref_ptee;
2343
2344 if (mapTypeMod == "is_device_ptr")
2345 mapTypeBits |= ClauseMapFlags::is_device_ptr;
2346
2347 if (mapTypeMod == "target_param")
2348 mapTypeBits |= ClauseMapFlags::target_param;
2349
2350 return success();
2351 };
2352
2353 if (parser.parseCommaSeparatedList(parseTypeAndMod))
2354 return failure();
2355
2356 mapType =
2357 parser.getBuilder().getAttr<mlir::omp::ClauseMapFlagsAttr>(mapTypeBits);
2358
2359 return success();
2360}
2361
2362/// Prints a map_entries map type from its numeric value out into its string
2363/// format.
2364static void printMapClause(OpAsmPrinter &p, Operation *op,
2365 ClauseMapFlagsAttr mapType) {
2367 ClauseMapFlags mapFlags = mapType.getValue();
2368
2369 // handling of always, close, present placed at the beginning of the string
2370 // to aid readability
2371 if (mapTypeToBool(mapFlags, ClauseMapFlags::always))
2372 mapTypeStrs.push_back("always");
2373 if (mapTypeToBool(mapFlags, ClauseMapFlags::implicit))
2374 mapTypeStrs.push_back("implicit");
2375 if (mapTypeToBool(mapFlags, ClauseMapFlags::ompx_hold))
2376 mapTypeStrs.push_back("ompx_hold");
2377 if (mapTypeToBool(mapFlags, ClauseMapFlags::close))
2378 mapTypeStrs.push_back("close");
2379 if (mapTypeToBool(mapFlags, ClauseMapFlags::present))
2380 mapTypeStrs.push_back("present");
2381 if (mapTypeToBool(mapFlags, ClauseMapFlags::target_param))
2382 mapTypeStrs.push_back("target_param");
2383
2384 // special handling of to/from/tofrom/delete and release/alloc, release +
2385 // alloc are the abscense of one of the other flags, whereas tofrom requires
2386 // both the to and from flag to be set.
2387 bool to = mapTypeToBool(mapFlags, ClauseMapFlags::to);
2388 bool from = mapTypeToBool(mapFlags, ClauseMapFlags::from);
2389
2390 if (to && from)
2391 mapTypeStrs.push_back("tofrom");
2392 else if (from)
2393 mapTypeStrs.push_back("from");
2394 else if (to)
2395 mapTypeStrs.push_back("to");
2396
2397 if (mapTypeToBool(mapFlags, ClauseMapFlags::del))
2398 mapTypeStrs.push_back("delete");
2399 if (mapTypeToBool(mapFlags, ClauseMapFlags::return_param))
2400 mapTypeStrs.push_back("return_param");
2401 if (mapTypeToBool(mapFlags, ClauseMapFlags::storage))
2402 mapTypeStrs.push_back("storage");
2403 if (mapTypeToBool(mapFlags, ClauseMapFlags::priv))
2404 mapTypeStrs.push_back("private");
2405 if (mapTypeToBool(mapFlags, ClauseMapFlags::literal))
2406 mapTypeStrs.push_back("literal");
2407 if (mapTypeToBool(mapFlags, ClauseMapFlags::attach))
2408 mapTypeStrs.push_back("attach");
2409 if (mapTypeToBool(mapFlags, ClauseMapFlags::attach_always))
2410 mapTypeStrs.push_back("attach_always");
2411 if (mapTypeToBool(mapFlags, ClauseMapFlags::attach_never))
2412 mapTypeStrs.push_back("attach_never");
2413 if (mapTypeToBool(mapFlags, ClauseMapFlags::attach_auto))
2414 mapTypeStrs.push_back("attach_auto");
2415 if (mapTypeToBool(mapFlags, ClauseMapFlags::ref_ptr))
2416 mapTypeStrs.push_back("ref_ptr");
2417 if (mapTypeToBool(mapFlags, ClauseMapFlags::ref_ptee))
2418 mapTypeStrs.push_back("ref_ptee");
2419 if (mapTypeToBool(mapFlags, ClauseMapFlags::is_device_ptr))
2420 mapTypeStrs.push_back("is_device_ptr");
2421 if (mapFlags == ClauseMapFlags::none)
2422 mapTypeStrs.push_back("none");
2423
2424 for (unsigned int i = 0; i < mapTypeStrs.size(); ++i) {
2425 p << mapTypeStrs[i];
2426 if (i + 1 < mapTypeStrs.size()) {
2427 p << ", ";
2428 }
2429 }
2430}
2431
2432static ParseResult parseMembersIndex(OpAsmParser &parser,
2433 ArrayAttr &membersIdx) {
2434 SmallVector<Attribute> values, memberIdxs;
2435
2436 auto parseIndices = [&]() -> ParseResult {
2437 int64_t value;
2438 if (parser.parseInteger(value))
2439 return failure();
2440 values.push_back(IntegerAttr::get(parser.getBuilder().getIntegerType(64),
2441 APInt(64, value, /*isSigned=*/false)));
2442 return success();
2443 };
2444
2445 do {
2446 if (failed(parser.parseLSquare()))
2447 return failure();
2448
2449 if (parser.parseCommaSeparatedList(parseIndices))
2450 return failure();
2451
2452 if (failed(parser.parseRSquare()))
2453 return failure();
2454
2455 memberIdxs.push_back(ArrayAttr::get(parser.getContext(), values));
2456 values.clear();
2457 } while (succeeded(parser.parseOptionalComma()));
2458
2459 if (!memberIdxs.empty())
2460 membersIdx = ArrayAttr::get(parser.getContext(), memberIdxs);
2461
2462 return success();
2463}
2464
2465static void printMembersIndex(OpAsmPrinter &p, MapInfoOp op,
2466 ArrayAttr membersIdx) {
2467 if (!membersIdx)
2468 return;
2469
2470 llvm::interleaveComma(membersIdx, p, [&p](Attribute v) {
2471 p << "[";
2472 auto memberIdx = cast<ArrayAttr>(v);
2473 llvm::interleaveComma(memberIdx.getValue(), p, [&p](Attribute v2) {
2474 p << cast<IntegerAttr>(v2).getInt();
2475 });
2476 p << "]";
2477 });
2478}
2479
2481 VariableCaptureKindAttr mapCaptureType) {
2482 std::string typeCapStr;
2483 llvm::raw_string_ostream typeCap(typeCapStr);
2484 if (mapCaptureType.getValue() == mlir::omp::VariableCaptureKind::ByRef)
2485 typeCap << "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)
2491 typeCap << "This";
2492 p << typeCapStr;
2493}
2494
2495static ParseResult parseCaptureType(OpAsmParser &parser,
2496 VariableCaptureKindAttr &mapCaptureType) {
2497 StringRef mapCaptureKey;
2498 if (parser.parseKeyword(&mapCaptureKey))
2499 return failure();
2500
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);
2513
2514 return success();
2515}
2516
2517static LogicalResult verifyMapInfoForMapClause(
2518 Operation *op, mlir::omp::MapInfoOp mapInfoOp,
2521 &updateFromVars) {
2522 mlir::omp::ClauseMapFlags mapTypeBits = mapInfoOp.getMapType();
2523
2524 bool to = mapTypeToBool(mapTypeBits, ClauseMapFlags::to);
2525 bool from = mapTypeToBool(mapTypeBits, ClauseMapFlags::from);
2526 bool del = mapTypeToBool(mapTypeBits, ClauseMapFlags::del);
2527
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);
2532
2533 if ((isa<TargetDataOp>(op) || isa<TargetOp>(op)) && del)
2534 return emitError(op->getLoc(),
2535 "to, from, tofrom and alloc map types are permitted");
2536
2537 if (isa<TargetEnterDataOp>(op) && (from || del))
2538 return emitError(op->getLoc(), "to and alloc map types are permitted");
2539
2540 if (isa<TargetExitDataOp>(op) && to)
2541 return emitError(op->getLoc(),
2542 "from, release and delete map types are permitted");
2543
2544 if (isa<TargetUpdateOp>(op)) {
2545 if (del) {
2546 return emitError(op->getLoc(),
2547 "at least one of to or from map types must be "
2548 "specified, other map types are not permitted");
2549 }
2550
2551 if (!to && !from && !attach) {
2552 return emitError(op->getLoc(),
2553 "at least one of to or from or attach map types must be "
2554 "specified, other map types are not permitted");
2555 }
2556
2557 auto updateVar = mapInfoOp.getVarPtr();
2558
2559 if ((to && from) || (to && updateFromVars.contains(updateVar)) ||
2560 (from && updateToVars.contains(updateVar))) {
2561 return emitError(
2562 op->getLoc(),
2563 "either to or from map types can be specified, not both");
2564 }
2565
2566 if (always || close || implicit) {
2567 return emitError(
2568 op->getLoc(),
2569 "present, mapper and iterator map type modifiers are permitted");
2570 }
2571
2572 // It's possible we have an attach map, in which case if there is no to
2573 // or from tied to it, we skip insertion.
2574 if (to || from) {
2575 to ? updateToVars.insert(updateVar) : updateFromVars.insert(updateVar);
2576 }
2577 }
2578
2579 if ((mapInfoOp.getVarPtrPtr() && !mapInfoOp.getVarPtrPtrType()) ||
2580 (!mapInfoOp.getVarPtrPtr() && mapInfoOp.getVarPtrPtrType())) {
2581 return emitError(op->getLoc(),
2582 "if varPtrPtr or varPtrPtrType is specified, then both "
2583 "must be present");
2584 }
2585
2586 return success();
2587}
2588
2589static LogicalResult verifyMapClause(Operation *op, OperandRange mapVars,
2590 OperandRange mapIterated) {
2593
2594 for (auto mapOp : mapVars) {
2595 if (!mapOp.getDefiningOp())
2596 return emitError(op->getLoc(), "missing map operation");
2597
2598 if (auto mapInfoOp = mapOp.getDefiningOp<mlir::omp::MapInfoOp>()) {
2599 if (failed(verifyMapInfoForMapClause(op, mapInfoOp, updateToVars,
2600 updateFromVars)))
2601 return failure();
2602 } else if (!isa<DeclareMapperInfoOp>(op)) {
2603 return emitError(op->getLoc(),
2604 "map argument is not a map entry operation");
2605 }
2606 }
2607
2608 // Verify iterated map entries.
2609 for (auto iterVal : mapIterated) {
2610 auto iterOp = iterVal.getDefiningOp<mlir::omp::IteratorOp>();
2611 if (!iterOp)
2612 return op->emitOpError() << "'map_iterated' arguments must be defined by "
2613 "'omp.iterator' ops";
2614
2615 // Check that the iterator body yields a value defined by omp.map.info.
2616 auto yieldOp =
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'";
2623
2624 if (failed(verifyMapInfoForMapClause(op, yieldedMapInfo, updateToVars,
2625 updateFromVars)))
2626 return failure();
2627 }
2628
2629 return success();
2630}
2631
2632template <typename OpType>
2633static LogicalResult verifyPrivateVarList(OpType &op);
2634
2635static LogicalResult verifyPrivateVarsMapping(TargetOp targetOp) {
2636 std::optional<DenseI64ArrayAttr> privateMapIndices =
2637 targetOp.getPrivateMapsAttr();
2638
2639 // None of the private operands are mapped.
2640 if (!privateMapIndices.has_value() || !privateMapIndices.value())
2641 return success();
2642
2643 OperandRange privateVars = targetOp.getPrivateVars();
2644
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");
2649
2650 return success();
2651}
2652
2653//===----------------------------------------------------------------------===//
2654// MapInfoOp
2655//===----------------------------------------------------------------------===//
2656
2657static LogicalResult verifyMapInfoDefinedArgs(Operation *op,
2658 StringRef clauseName,
2659 OperandRange vars) {
2660 for (Value var : vars)
2661 if (!llvm::isa_and_present<MapInfoOp>(var.getDefiningOp()))
2662 return op->emitOpError()
2663 << "'" << clauseName
2664 << "' arguments must be defined by 'omp.map.info' ops";
2665 return success();
2666}
2667
2668LogicalResult MapInfoOp::verify() {
2669 if (getMapperId() &&
2671 *this, getMapperIdAttr())) {
2672 return emitError("invalid mapper id");
2673 }
2674
2675 if (failed(verifyMapInfoDefinedArgs(*this, "members", getMembers())))
2676 return failure();
2677
2678 return success();
2679}
2680
2681//===----------------------------------------------------------------------===//
2682// TargetDataOp
2683//===----------------------------------------------------------------------===//
2684
2685void TargetDataOp::build(OpBuilder &builder, OperationState &state,
2686 const TargetDataOperands &clauses) {
2687 TargetDataOp::build(builder, state, clauses.device, clauses.ifExpr,
2688 clauses.mapVars, clauses.mapIterated,
2689 clauses.useDeviceAddrVars, clauses.useDevicePtrVars);
2690}
2691
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");
2698 }
2699
2700 if (failed(verifyMapInfoDefinedArgs(*this, "use_device_ptr",
2701 getUseDevicePtrVars())))
2702 return failure();
2703
2704 if (failed(verifyMapInfoDefinedArgs(*this, "use_device_addr",
2705 getUseDeviceAddrVars())))
2706 return failure();
2707
2708 return verifyMapClause(*this, getMapVars(), getMapIterated());
2709}
2710
2711//===----------------------------------------------------------------------===//
2712// TargetEnterDataOp
2713//===----------------------------------------------------------------------===//
2714
2715void TargetEnterDataOp::build(
2716 OpBuilder &builder, OperationState &state,
2717 const TargetEnterExitUpdateDataOperands &clauses) {
2718 MLIRContext *ctx = builder.getContext();
2719 TargetEnterDataOp::build(
2720 builder, state, makeArrayAttr(ctx, clauses.dependKinds),
2721 clauses.dependVars, makeArrayAttr(ctx, clauses.dependIteratedKinds),
2722 clauses.dependIterated, clauses.device, clauses.ifExpr, clauses.mapVars,
2723 clauses.mapIterated, clauses.nowait);
2724}
2725
2726LogicalResult TargetEnterDataOp::verify() {
2727 LogicalResult verifyDependVars =
2728 verifyDependVarList(*this, getDependKinds(), getDependVars(),
2729 getDependIteratedKinds(), getDependIterated());
2730 return failed(verifyDependVars)
2731 ? verifyDependVars
2732 : verifyMapClause(*this, getMapVars(), getMapIterated());
2733}
2734
2735//===----------------------------------------------------------------------===//
2736// TargetExitDataOp
2737//===----------------------------------------------------------------------===//
2738
2739void TargetExitDataOp::build(OpBuilder &builder, OperationState &state,
2740 const TargetEnterExitUpdateDataOperands &clauses) {
2741 MLIRContext *ctx = builder.getContext();
2742 TargetExitDataOp::build(
2743 builder, state, makeArrayAttr(ctx, clauses.dependKinds),
2744 clauses.dependVars, makeArrayAttr(ctx, clauses.dependIteratedKinds),
2745 clauses.dependIterated, clauses.device, clauses.ifExpr, clauses.mapVars,
2746 clauses.mapIterated, clauses.nowait);
2747}
2748
2749LogicalResult TargetExitDataOp::verify() {
2750 LogicalResult verifyDependVars =
2751 verifyDependVarList(*this, getDependKinds(), getDependVars(),
2752 getDependIteratedKinds(), getDependIterated());
2753 return failed(verifyDependVars)
2754 ? verifyDependVars
2755 : verifyMapClause(*this, getMapVars(), getMapIterated());
2756}
2757
2758//===----------------------------------------------------------------------===//
2759// TargetUpdateOp
2760//===----------------------------------------------------------------------===//
2761
2762void TargetUpdateOp::build(OpBuilder &builder, OperationState &state,
2763 const TargetEnterExitUpdateDataOperands &clauses) {
2764 MLIRContext *ctx = builder.getContext();
2765 TargetUpdateOp::build(builder, state, makeArrayAttr(ctx, clauses.dependKinds),
2766 clauses.dependVars,
2767 makeArrayAttr(ctx, clauses.dependIteratedKinds),
2768 clauses.dependIterated, clauses.device, clauses.ifExpr,
2769 clauses.mapVars, clauses.mapIterated, clauses.nowait);
2770}
2771
2772LogicalResult TargetUpdateOp::verify() {
2773 LogicalResult verifyDependVars =
2774 verifyDependVarList(*this, getDependKinds(), getDependVars(),
2775 getDependIteratedKinds(), getDependIterated());
2776 return failed(verifyDependVars)
2777 ? verifyDependVars
2778 : verifyMapClause(*this, getMapVars(), getMapIterated());
2779}
2780
2781//===----------------------------------------------------------------------===//
2782// TargetOp
2783//===----------------------------------------------------------------------===//
2784
2785void TargetOp::build(OpBuilder &builder, OperationState &state,
2786 const TargetExtOperands &clauses) {
2787 MLIRContext *ctx = builder.getContext();
2788 TargetOp::build(
2789 builder, state, clauses.allocateVars, clauses.allocatorVars,
2790 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
2791 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
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,
2798 makeDenseBoolArrayAttr(ctx, clauses.inReductionByref),
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, /*private_maps=*/nullptr, clauses.kernelType);
2803}
2804
2805bool TargetOp::hasHostEvalTripCount() {
2806 TargetExecMode mode = getKernelType();
2807 if (mode == TargetExecMode::spmd || mode == TargetExecMode::spmd_no_loop)
2808 return true;
2809
2810 if (mode == TargetExecMode::bare)
2811 return false;
2812
2813 // If it represents a `target teams distribute` construct, also evaluate the
2814 // `distribute` trip count on the host.
2815 Operation *capturedOp =
2816 cast<ComposableOpInterface>(getOperation()).findCapturedOp();
2817 if (auto loopNestOp = dyn_cast_if_present<LoopNestOp>(capturedOp)) {
2819 loopNestOp.gatherWrappers(loopWrappers);
2820
2821 LoopWrapperInterface *innermostWrapper = loopWrappers.begin();
2822 if (isa<SimdOp>(innermostWrapper))
2823 innermostWrapper = std::next(innermostWrapper);
2824
2825 auto numWrappers = std::distance(innermostWrapper, loopWrappers.end());
2826 if (numWrappers != 1)
2827 return false;
2828
2829 if (!isa<DistributeOp>(innermostWrapper))
2830 return false;
2831
2832 Operation *parentOp = innermostWrapper->getOperation()->getParentOp();
2833 if (isa_and_present<TeamsOp>(parentOp) &&
2834 parentOp->getParentOp() == getOperation())
2835 return true;
2836 }
2837
2838 return false;
2839}
2840
2841/// An `omp.target` `in_reduction` operand is captured by a `map_entries` entry
2842/// when the entry's `MapInfoOp` var_ptr is the same SSA value, or another
2843/// result of the same defining op. At this stage, exact identity can only be
2844/// required for block arguments, which have no defining op. Flang emits
2845/// `hlfir.declare` #0 for the `in_reduction` operand and #1 for the map
2846/// `var_ptr`; these collapse to the same value after lowering, but that cannot
2847/// be enforced here.
2848static bool targetInReductionCapturedBy(Value inReductionVar, Value mapVarPtr) {
2849 if (mapVarPtr == inReductionVar)
2850 return true;
2851 Operation *def = inReductionVar.getDefiningOp();
2852 return def && mapVarPtr.getDefiningOp() == def;
2853}
2854
2855LogicalResult TargetOp::verify() {
2857 getOperation(), getAllocateVars(), getAllocatorVars(),
2858 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
2859 getPrivateVars(), getPrivateSymsAttr())))
2860 return failure();
2861
2862 if (getKernelType() == TargetExecMode::bare && !isCombined())
2863 return emitOpError() << "bare kernel requires 'omp.combined'";
2864
2865 if (failed(verifyDependVarList(*this, getDependKinds(), getDependVars(),
2866 getDependIteratedKinds(),
2867 getDependIterated())))
2868 return failure();
2869
2870 if (failed(verifyMapInfoDefinedArgs(*this, "has_device_addr",
2871 getHasDeviceAddrVars())))
2872 return failure();
2873
2874 if (failed(verifyMapClause(*this, getMapVars(), getMapIterated())))
2875 return failure();
2876
2878 *this, getDynGroupprivateAccessGroupAttr(),
2879 getDynGroupprivateFallbackAttr(), getDynGroupprivateSize())))
2880 return failure();
2881
2882 if (failed(verifyPrivateVarList(*this)))
2883 return failure();
2884
2885 if (failed(verifyReductionVarList(*this, getInReductionSyms(),
2886 getInReductionVars(),
2887 getInReductionByref())))
2888 return failure();
2889
2890 // An `in_reduction` operand on `omp.target` has no dedicated entry block
2891 // argument; inside the region it is accessed through the block argument of a
2892 // matching `map_entries` entry, and the host rewrites that map argument to
2893 // the reduction-private storage. Require every `in_reduction` operand to be
2894 // captured by at least one `map_entries` entry.
2895 for (Value inReductionVar : getInReductionVars()) {
2896 bool captured = false;
2897 for (Value mapVar : getMapVars()) {
2898 auto mapInfo = mapVar.getDefiningOp<MapInfoOp>();
2899 if (targetInReductionCapturedBy(inReductionVar, mapInfo.getVarPtr())) {
2900 captured = true;
2901 break;
2902 }
2903 }
2904 if (!captured)
2905 return emitOpError() << "in_reduction variable must be captured by a "
2906 "matching map_entries entry";
2907 }
2908
2909 return verifyPrivateVarsMapping(*this);
2910}
2911
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");
2917
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";
2926 default:
2927 break;
2928 }
2929 }
2930
2931 Operation *capturedOp =
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";
2938
2939 bool isTargetDevice = false;
2940 if (auto offloadMod = (*this)->getParentOfType<OffloadModuleInterface>())
2941 if (offloadMod.getIsTargetDevice())
2942 isTargetDevice = true;
2943
2944 // Check that host_eval values are only used in legal ways.
2945 llvm::ArrayRef<BlockArgument> hostEvalBlockArgs =
2946 cast<BlockArgOpenMPOpInterface>(getOperation()).getHostEvalBlockArgs();
2947
2948 bool hostEvalTripCount = hasHostEvalTripCount();
2949 for (Value hostEvalArg : hostEvalBlockArgs) {
2950 for (Operation *user : hostEvalArg.getUsers()) {
2951 if (auto teamsOp = dyn_cast<TeamsOp>(user)) {
2952 // Check if used in num_teams_lower or any of num_teams_upper_vars
2953 if (hostEvalArg == teamsOp.getNumTeamsLower() ||
2954 llvm::is_contained(teamsOp.getNumTeamsUpperVars(), hostEvalArg) ||
2955 llvm::is_contained(teamsOp.getThreadLimitVars(), hostEvalArg))
2956 continue;
2957
2958 return emitOpError() << "host_eval argument only legal as 'num_teams' "
2959 "and 'thread_limit' in 'omp.teams'";
2960 }
2961 if (auto parallelOp = dyn_cast<ParallelOp>(user)) {
2962 if (llvm::is_contained(parallelOp.getNumThreadsVars(), hostEvalArg))
2963 continue;
2964
2965 return emitOpError()
2966 << "host_eval argument only legal as 'num_threads' in "
2967 "'omp.parallel'";
2968 }
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)))
2974 continue;
2975
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";
2979 }
2980
2981 return emitOpError() << "host_eval argument illegal use in '"
2982 << user->getName() << "' operation";
2983 }
2984 }
2985
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";
2994 }
2995 }
2996
2997 return success();
2998}
2999
3000//===----------------------------------------------------------------------===//
3001// ParallelOp
3002//===----------------------------------------------------------------------===//
3003
3004void ParallelOp::build(OpBuilder &builder, OperationState &state,
3005 ArrayRef<NamedAttribute> attributes) {
3006 ParallelOp::build(builder, state, /*allocate_vars=*/ValueRange(),
3007 /*allocator_vars=*/ValueRange(),
3008 /*allocate_alignments=*/nullptr,
3009 /*allocate_private_indices=*/nullptr, /*if_expr=*/nullptr,
3010 /*num_threads_vars=*/ValueRange(),
3011 /*private_vars=*/ValueRange(),
3012 /*private_syms=*/nullptr, /*private_needs_barrier=*/false,
3013 /*proc_bind_kind=*/nullptr,
3014 /*reduction_mod =*/nullptr, /*reduction_vars=*/ValueRange(),
3015 /*reduction_byref=*/nullptr, /*reduction_syms=*/nullptr);
3016 state.addAttributes(attributes);
3017}
3018
3019void ParallelOp::build(OpBuilder &builder, OperationState &state,
3020 const ParallelOperands &clauses) {
3021 MLIRContext *ctx = builder.getContext();
3022 ParallelOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3023 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3024 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3025 clauses.ifExpr, clauses.numThreadsVars, clauses.privateVars,
3026 makeArrayAttr(ctx, clauses.privateSyms),
3027 clauses.privateNeedsBarrier, clauses.procBindKind,
3028 clauses.reductionMod, clauses.reductionVars,
3029 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3030 makeArrayAttr(ctx, clauses.reductionSyms));
3031}
3032
3033template <typename OpType>
3034static LogicalResult verifyPrivateVarList(OpType &op) {
3035 auto privateVars = op.getPrivateVars();
3036 auto privateSyms = op.getPrivateSymsAttr();
3037
3038 if (privateVars.empty() && (privateSyms == nullptr || privateSyms.empty()))
3039 return success();
3040
3041 auto numPrivateVars = privateVars.size();
3042 auto numPrivateSyms = (privateSyms == nullptr) ? 0 : privateSyms.size();
3043
3044 if (numPrivateVars != numPrivateSyms)
3045 return op.emitError() << "inconsistent number of private variables and "
3046 "privatizer op symbols, private vars: "
3047 << numPrivateVars
3048 << " vs. privatizer op symbols: " << numPrivateSyms;
3049
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 =
3055
3056 if (privatizerOp == nullptr)
3057 return op.emitError() << "failed to lookup privatizer op with symbol: '"
3058 << privateSym << "'";
3059
3060 Type privatizerType = privatizerOp.getArgType();
3061
3062 if (privatizerType && (varType != privatizerType))
3063 return op.emitError()
3064 << "type mismatch between a "
3065 << (privatizerOp.getDataSharingType() ==
3066 DataSharingClauseType::Private
3067 ? "private"
3068 : "firstprivate")
3069 << " variable and its privatizer op, var type: " << varType
3070 << " vs. privatizer op type: " << privatizerType;
3071 }
3072
3073 return success();
3074}
3075
3076LogicalResult ParallelOp::verify() {
3077 if (failed(verifyPrivateVarList(*this)))
3078 return failure();
3080 getOperation(), getAllocateVars(), getAllocatorVars(),
3081 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3082 getPrivateVars(), getPrivateSymsAttr(),
3083 /*requirePrivateIndices=*/true)))
3084 return failure();
3085
3086 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3087 getReductionByref());
3088}
3089
3090LogicalResult ParallelOp::verifyRegions() {
3091 auto distChildOps = getOps<DistributeOp>();
3092 int numDistChildOps = std::distance(distChildOps.begin(), distChildOps.end());
3093 if (numDistChildOps > 1)
3094 return emitError()
3095 << "multiple 'omp.distribute' nested inside of 'omp.parallel'";
3096
3097 if (numDistChildOps == 1) {
3098 if (!isComposite())
3099 return emitError()
3100 << "'omp.composite' attribute missing from composite operation";
3101
3102 auto *ompDialect = getContext()->getLoadedDialect<OpenMPDialect>();
3103 Operation &distributeOp = **distChildOps.begin();
3104 for (Operation &childOp : getOps()) {
3105 if (&childOp == &distributeOp || ompDialect != childOp.getDialect())
3106 continue;
3107
3108 if (!childOp.hasTrait<OpTrait::IsTerminator>())
3109 return emitError() << "unexpected OpenMP operation inside of composite "
3110 "'omp.parallel': "
3111 << childOp.getName();
3112 }
3113 } else if (isComposite()) {
3114 return emitError()
3115 << "'omp.composite' attribute present in non-composite operation";
3116 }
3117 return success();
3118}
3119
3120//===----------------------------------------------------------------------===//
3121// TeamsOp
3122//===----------------------------------------------------------------------===//
3123
3125 while ((op = op->getParentOp()))
3126 if (isa<OpenMPDialect>(op->getDialect()))
3127 return false;
3128 return true;
3129}
3130
3131void TeamsOp::build(OpBuilder &builder, OperationState &state,
3132 const TeamsOperands &clauses) {
3133 MLIRContext *ctx = builder.getContext();
3134 // TODO Store clauses in op: privateVars, privateSyms, privateNeedsBarrier
3135 TeamsOp::build(
3136 builder, state, clauses.allocateVars, clauses.allocatorVars,
3137 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3138 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3139 clauses.dynGroupprivateAccessGroup, clauses.dynGroupprivateFallback,
3140 clauses.dynGroupprivateSize, clauses.ifExpr, clauses.numTeamsLower,
3141 clauses.numTeamsUpperVars, /*private_vars=*/{}, /*private_syms=*/nullptr,
3142 /*private_needs_barrier=*/false, clauses.reductionMod,
3143 clauses.reductionVars,
3144 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3145 makeArrayAttr(ctx, clauses.reductionSyms), clauses.threadLimitVars);
3146}
3147
3148// Verify num_teams clause
3149static LogicalResult verifyNumTeamsClause(Operation *op, Value numTeamsLower,
3150 OperandRange numTeamsUpperVars) {
3151 // If lower is specified, upper must have exactly one value
3152 if (numTeamsLower) {
3153 if (numTeamsUpperVars.size() != 1)
3154 return op->emitError(
3155 "expected exactly one num_teams upper bound when lower bound is "
3156 "specified");
3157 if (numTeamsLower.getType() != numTeamsUpperVars[0].getType())
3158 return op->emitError(
3159 "expected num_teams upper bound and lower bound to be "
3160 "the same type");
3161 }
3162
3163 return success();
3164}
3165
3166LogicalResult TeamsOp::verify() {
3167 // Check parent region
3168 // TODO If nested inside of a target region, also check that it does not
3169 // contain any statements, declarations or directives other than this
3170 // omp.teams construct. The issue is how to support the initialization of
3171 // this operation's own arguments (allow SSA values across omp.target?).
3172 Operation *op = getOperation();
3173 auto parentTarget = llvm::dyn_cast_if_present<TargetOp>(op->getParentOp());
3174 if (!parentTarget && !opInGlobalImplicitParallelRegion(op))
3175 return emitError("expected to be nested inside of omp.target or not nested "
3176 "in any OpenMP dialect operations");
3177
3178 // Check for num_teams clause restrictions
3179 if (failed(verifyNumTeamsClause(op, this->getNumTeamsLower(),
3180 this->getNumTeamsUpperVars())))
3181 return failure();
3182
3183 if (parentTarget &&
3184 parentTarget.getKernelType() == TargetExecMode::spmd_no_loop &&
3185 (getNumTeamsLower() || !getNumTeamsUpperVars().empty()))
3186 return emitOpError() << "'num_teams' not allowed in SPMD-no-loop kernels";
3187
3189 getOperation(), getAllocateVars(), getAllocatorVars(),
3190 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3191 getPrivateVars(), getPrivateSymsAttr())))
3192 return failure();
3193
3195 op, getDynGroupprivateAccessGroupAttr(),
3196 getDynGroupprivateFallbackAttr(), getDynGroupprivateSize())))
3197 return failure();
3198
3199 if (failed(verifyPrivateVarList(*this)))
3200 return failure();
3201
3202 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3203 getReductionByref());
3204}
3205
3206//===----------------------------------------------------------------------===//
3207// SectionOp
3208//===----------------------------------------------------------------------===//
3209
3210OperandRange SectionOp::getPrivateVars() {
3211 return getParentOp().getPrivateVars();
3212}
3213
3214OperandRange SectionOp::getReductionVars() {
3215 return getParentOp().getReductionVars();
3216}
3217
3218//===----------------------------------------------------------------------===//
3219// SectionsOp
3220//===----------------------------------------------------------------------===//
3221
3222void SectionsOp::build(OpBuilder &builder, OperationState &state,
3223 const SectionsOperands &clauses) {
3224 MLIRContext *ctx = builder.getContext();
3225 // TODO Store clauses in op: privateVars, privateSyms, privateNeedsBarrier
3226 SectionsOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3227 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3228 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3229 clauses.nowait, /*private_vars=*/{},
3230 /*private_syms=*/nullptr, /*private_needs_barrier=*/nullptr,
3231 clauses.reductionMod, clauses.reductionVars,
3232 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3233 makeArrayAttr(ctx, clauses.reductionSyms));
3234}
3235
3236LogicalResult SectionsOp::verify() {
3237 if (isCombined())
3238 return emitOpError() << "cannot be a non-innermost combined construct leaf";
3239
3241 getOperation(), getAllocateVars(), getAllocatorVars(),
3242 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3243 getPrivateVars(), getPrivateSymsAttr())))
3244 return failure();
3245
3246 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3247 getReductionByref());
3248}
3249
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";
3255 }
3256 }
3257
3258 return success();
3259}
3260
3261//===----------------------------------------------------------------------===//
3262// ScopeOp
3263//===----------------------------------------------------------------------===//
3264
3265void ScopeOp::build(OpBuilder &builder, OperationState &state,
3266 const ScopeOperands &clauses) {
3267 MLIRContext *ctx = builder.getContext();
3268 ScopeOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3269 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3270 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3271 clauses.nowait, clauses.privateVars,
3272 makeArrayAttr(ctx, clauses.privateSyms),
3273 clauses.privateNeedsBarrier, clauses.reductionMod,
3274 clauses.reductionVars,
3275 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3276 makeArrayAttr(ctx, clauses.reductionSyms));
3277}
3278
3279LogicalResult ScopeOp::verify() {
3281 getOperation(), getAllocateVars(), getAllocatorVars(),
3282 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3283 getPrivateVars(), getPrivateSymsAttr(),
3284 /*requirePrivateIndices=*/true)))
3285 return failure();
3286
3287 if (failed(verifyPrivateVarList(*this)))
3288 return failure();
3289
3290 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3291 getReductionByref());
3292}
3293
3294//===----------------------------------------------------------------------===//
3295// SingleOp
3296//===----------------------------------------------------------------------===//
3297
3298void SingleOp::build(OpBuilder &builder, OperationState &state,
3299 const SingleOperands &clauses) {
3300 MLIRContext *ctx = builder.getContext();
3301 // TODO Store clauses in op: privateVars, privateSyms, privateNeedsBarrier
3302 SingleOp::build(builder, state, clauses.allocateVars, clauses.allocatorVars,
3303 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3304 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3305 clauses.copyprivateVars,
3306 makeArrayAttr(ctx, clauses.copyprivateSyms), clauses.nowait,
3307 /*private_vars=*/{}, /*private_syms=*/nullptr,
3308 /*private_needs_barrier=*/nullptr);
3309}
3310
3311LogicalResult SingleOp::verify() {
3313 getOperation(), getAllocateVars(), getAllocatorVars(),
3314 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3315 getPrivateVars(), getPrivateSymsAttr())))
3316 return failure();
3317
3318 return verifyCopyprivateVarList(*this, getCopyprivateVars(),
3319 getCopyprivateSyms());
3320}
3321
3322//===----------------------------------------------------------------------===//
3323// WorkshareOp
3324//===----------------------------------------------------------------------===//
3325
3326void WorkshareOp::build(OpBuilder &builder, OperationState &state,
3327 const WorkshareOperands &clauses) {
3328 WorkshareOp::build(builder, state, clauses.nowait);
3329}
3330
3331LogicalResult WorkshareOp::verify() {
3332 if (isCombined())
3333 return emitOpError() << "cannot be a non-innermost combined construct leaf";
3334
3335 return success();
3336}
3337
3338//===----------------------------------------------------------------------===//
3339// WorkshareLoopWrapperOp
3340//===----------------------------------------------------------------------===//
3341
3342LogicalResult WorkshareLoopWrapperOp::verifyRegions() {
3343 if (isa_and_nonnull<LoopWrapperInterface>((*this)->getParentOp()) ||
3344 getNestedWrapper())
3345 return emitOpError() << "expected to be a standalone loop wrapper";
3346
3347 return success();
3348}
3349
3350//===----------------------------------------------------------------------===//
3351// LoopWrapperInterface
3352//===----------------------------------------------------------------------===//
3353
3354LogicalResult LoopWrapperInterface::verifyImpl() {
3355 Operation *op = this->getOperation();
3356 if (!op->hasTrait<OpTrait::NoTerminator>() ||
3358 return emitOpError() << "loop wrapper must also have the `NoTerminator` "
3359 "and `SingleBlock` traits";
3360
3361 if (op->getNumRegions() != 1)
3362 return emitOpError() << "loop wrapper does not contain exactly one region";
3363
3364 Region &region = op->getRegion(0);
3365 if (range_size(region.getOps()) != 1)
3366 return emitOpError()
3367 << "loop wrapper does not contain exactly one nested op";
3368
3369 Operation &firstOp = *region.op_begin();
3370 if (!isa<LoopNestOp, LoopWrapperInterface>(firstOp))
3371 return emitOpError() << "nested in loop wrapper is not another loop "
3372 "wrapper or `omp.loop_nest`";
3373
3374 return success();
3375}
3376
3377//===----------------------------------------------------------------------===//
3378// ComposableOpInterface
3379//===----------------------------------------------------------------------===//
3380
3381Operation *ComposableOpInterface::findCapturedOp() {
3382 Operation *op = this->getOperation();
3383
3384 // Handle the composite case by returning the wrapped omp.loop_nest.
3385 if (auto wrapperOp = dyn_cast<LoopWrapperInterface>(op))
3386 return wrapperOp.getWrappedLoop();
3387
3388 // Do not look further if this op is not combined with any of its children.
3389 // Need to check for composite for the omp.parallel case, which is not a loop
3390 // wrapper itself.
3391 if (!isCombined() && !isComposite())
3392 return op;
3393
3394 Region &region = op->getRegion(0);
3395 for (Operation &nestedOp : region.getOps()) {
3396 if (auto wrapperOp = dyn_cast<LoopWrapperInterface>(&nestedOp))
3397 return wrapperOp.getWrappedLoop();
3398
3399 if (auto composableOp = dyn_cast<ComposableOpInterface>(&nestedOp))
3400 return composableOp.findCapturedOp();
3401 }
3402
3403 // This can only be reached if the op has an omp.combined attribute but the
3404 // corresponding nested composable op has been deleted. In that case, it's
3405 // correct to return this operation.
3406 return op;
3407}
3408
3409LogicalResult ComposableOpInterface::verifyImpl() {
3410 Operation *op = this->getOperation();
3411
3412 if (op->getNumRegions() != 1)
3413 return emitOpError() << "composable ops must have a single region";
3414
3415 if (isComposite() && !isa<LoopWrapperInterface, ParallelOp>(op))
3416 return emitOpError() << "non-loop wrapper cannot be composite";
3417
3418 // If combined, must have exactly one eligible nested op (composable or loop
3419 // wrapper).
3420 if (isCombined()) {
3421 Operation *nestedOp = nullptr;
3422 auto count = llvm::count_if(
3423 op->getRegion(0).getOps(), [&nestedOp](mlir::Operation &op) {
3424 if (isa<ComposableOpInterface, LoopWrapperInterface>(op)) {
3425 nestedOp = &op;
3426 return true;
3427 }
3428 return false;
3429 });
3430
3431 // Make an exception for ops marked as omp.combined with no eligible nested
3432 // ops: this situation should be disallowed, but it can be reached if an
3433 // MLIR optimization pass find that the child operation has no side effects
3434 // (many ComposableOpInterface ops have RecursiveMemoryEffects), so it gets
3435 // deleted without updating the parent's attribute.
3436 //
3437 // Since there's a well defined way of handling that situation (treat it as
3438 // non-combined), we relax the requirement here. Ensuring the parent is
3439 // updated every time a pass that can potentially remove a child composable
3440 // op runs is less preferable as a solution.
3441 if (count == 0)
3442 return success();
3443
3444 if (count > 1)
3445 return emitOpError()
3446 << "multiple eligible child ops found in combined op";
3447
3448 // This operation cannot be combined if its captured nested op can be
3449 // executed more than once (i.e. its block's successors can reach it) or if
3450 // it's not guaranteed to be executed before all exits of the region (i.e.
3451 // it doesn't dominate all blocks with no successors reachable from the
3452 // entry block).
3453 DominanceInfo domInfo;
3454 Block *parentBlock = nestedOp->getBlock();
3455
3456 for (Block *successor : parentBlock->getSuccessors())
3457 if (successor->isReachable(parentBlock))
3458 return emitOpError() << "nested combined child op is part of a loop";
3459
3460 for (Block &block : op->getRegion(0))
3461 if (domInfo.isReachableFromEntry(&block) && block.hasNoSuccessors() &&
3462 !domInfo.dominates(parentBlock, &block))
3463 return emitOpError()
3464 << "nested combined child op doesn't unconditionally execute";
3465 }
3466 return success();
3467}
3468
3469//===----------------------------------------------------------------------===//
3470// LoopOp
3471//===----------------------------------------------------------------------===//
3472
3473void LoopOp::build(OpBuilder &builder, OperationState &state,
3474 const LoopOperands &clauses) {
3475 MLIRContext *ctx = builder.getContext();
3476
3477 LoopOp::build(builder, state, clauses.bindKind, clauses.privateVars,
3478 makeArrayAttr(ctx, clauses.privateSyms),
3479 clauses.privateNeedsBarrier, clauses.order, clauses.orderMod,
3480 clauses.reductionMod, clauses.reductionVars,
3481 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3482 makeArrayAttr(ctx, clauses.reductionSyms));
3483}
3484
3485LogicalResult LoopOp::verify() {
3486 if (failed(verifyPrivateVarList(*this)))
3487 return failure();
3488
3489 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3490 getReductionByref());
3491}
3492
3493LogicalResult LoopOp::verifyRegions() {
3494 if (llvm::isa_and_nonnull<LoopWrapperInterface>((*this)->getParentOp()) ||
3495 getNestedWrapper())
3496 return emitOpError() << "expected to be a standalone loop wrapper";
3497
3498 return success();
3499}
3500
3501//===----------------------------------------------------------------------===//
3502// WsloopOp
3503//===----------------------------------------------------------------------===//
3504
3505void WsloopOp::build(OpBuilder &builder, OperationState &state,
3506 ArrayRef<NamedAttribute> attributes) {
3507 build(builder, state, /*allocate_vars=*/{}, /*allocator_vars=*/{},
3508 /*allocate_alignments=*/nullptr,
3509 /*allocate_private_indices=*/nullptr,
3510 /*linear_vars=*/ValueRange(), /*linear_step_vars=*/ValueRange(),
3511 /*linear_var_types*/ nullptr, /*linear_modifiers=*/nullptr,
3512 /*nowait=*/false, /*order=*/nullptr, /*order_mod=*/nullptr,
3513 /*ordered=*/nullptr, /*private_vars=*/{}, /*private_syms=*/nullptr,
3514 /*private_needs_barrier=*/false,
3515 /*reduction_mod=*/nullptr, /*reduction_vars=*/ValueRange(),
3516 /*reduction_byref=*/nullptr,
3517 /*reduction_syms=*/nullptr, /*schedule_kind=*/nullptr,
3518 /*schedule_chunk=*/nullptr, /*schedule_mod=*/nullptr,
3519 /*schedule_simd=*/false);
3520 state.addAttributes(attributes);
3521}
3522
3523void WsloopOp::build(OpBuilder &builder, OperationState &state,
3524 const WsloopOperands &clauses) {
3525 MLIRContext *ctx = builder.getContext();
3526 WsloopOp::build(
3527 builder, state, clauses.allocateVars, clauses.allocatorVars,
3528 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3529 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
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,
3535 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3536 makeArrayAttr(ctx, clauses.reductionSyms), clauses.scheduleKind,
3537 clauses.scheduleChunk, clauses.scheduleMod, clauses.scheduleSimd);
3538}
3539
3540LogicalResult WsloopOp::verify() {
3542 getOperation(), getAllocateVars(), getAllocatorVars(),
3543 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3544 getPrivateVars(), getPrivateSymsAttr())))
3545 return failure();
3546
3547 if (failed(
3548 verifyLinearModifiers(*this, getLinearModifiers(), getLinearVars())))
3549 return failure();
3550 if (getLinearVars().size() &&
3551 getLinearVarTypes().value().size() != getLinearVars().size())
3552 return emitError() << "Ill-formed type attributes for linear variables";
3553
3554 if (failed(verifyPrivateVarList(*this)))
3555 return failure();
3556
3557 return verifyReductionVarList(*this, getReductionSyms(), getReductionVars(),
3558 getReductionByref());
3559}
3560
3561LogicalResult WsloopOp::verifyRegions() {
3562 bool isCompositeChildLeaf =
3563 llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp());
3564
3565 if (LoopWrapperInterface nested = getNestedWrapper()) {
3566 if (!isComposite())
3567 return emitError()
3568 << "'omp.composite' attribute missing from composite wrapper";
3569
3570 // Check for the allowed leaf constructs that may appear in a composite
3571 // construct directly after DO/FOR.
3572 if (!isa<SimdOp>(nested))
3573 return emitError() << "only supported nested wrapper is 'omp.simd'";
3574
3575 } else if (isComposite() && !isCompositeChildLeaf) {
3576 return emitError()
3577 << "'omp.composite' attribute present in non-composite wrapper";
3578 } else if (!isComposite() && isCompositeChildLeaf) {
3579 return emitError()
3580 << "'omp.composite' attribute missing from composite wrapper";
3581 }
3582
3583 return success();
3584}
3585
3586//===----------------------------------------------------------------------===//
3587// Simd construct [2.9.3.1]
3588//===----------------------------------------------------------------------===//
3589
3590void SimdOp::build(OpBuilder &builder, OperationState &state,
3591 const SimdOperands &clauses) {
3592 MLIRContext *ctx = builder.getContext();
3593 SimdOp::build(builder, state, clauses.alignedVars,
3594 makeArrayAttr(ctx, clauses.alignments), clauses.ifExpr,
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,
3601 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3602 makeArrayAttr(ctx, clauses.reductionSyms), clauses.safelen,
3603 clauses.simdlen);
3604}
3605
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";
3612
3613 if (verifyAlignedClause(*this, getAlignments(), getAlignedVars()).failed())
3614 return failure();
3615
3616 if (verifyNontemporalClause(*this, getNontemporalVars()).failed())
3617 return failure();
3618
3619 if (failed(
3620 verifyLinearModifiers(*this, getLinearModifiers(), getLinearVars())))
3621 return failure();
3622
3623 bool isCompositeChildLeaf =
3624 llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp());
3625
3626 if (!isComposite() && isCompositeChildLeaf)
3627 return emitError()
3628 << "'omp.composite' attribute missing from composite wrapper";
3629
3630 if (isComposite() && !isCompositeChildLeaf)
3631 return emitError()
3632 << "'omp.composite' attribute present in non-composite wrapper";
3633
3634 // Firstprivate is not allowed for SIMD in the standard. Check that none of
3635 // the private decls are for firstprivate.
3636 std::optional<ArrayAttr> privateSyms = getPrivateSyms();
3637 if (privateSyms) {
3638 for (const Attribute &sym : *privateSyms) {
3639 auto symRef = cast<SymbolRefAttr>(sym);
3640 omp::PrivateClauseOp privatizer =
3642 getOperation(), symRef);
3643 if (!privatizer)
3644 return emitError() << "Cannot find privatizer '" << symRef << "'";
3645 if (privatizer.getDataSharingType() ==
3646 DataSharingClauseType::FirstPrivate)
3647 return emitError() << "FIRSTPRIVATE cannot be used with SIMD";
3648 }
3649 }
3650
3651 if (failed(verifyPrivateVarList(*this)))
3652 return failure();
3653
3654 if (getLinearVars().size() &&
3655 getLinearVarTypes().value().size() != getLinearVars().size())
3656 return emitError() << "Ill-formed type attributes for linear variables";
3657
3658 llvm::DenseSet<Value> privateVars(llvm::from_range, getPrivateVars());
3659 llvm::DenseSet<Value> reductionVars(llvm::from_range, getReductionVars());
3660 // TODO Check lastprivate vars when their support is added to SimdOp.
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";
3665 }
3666
3667 return success();
3668}
3669
3670LogicalResult SimdOp::verifyRegions() {
3671 if (getNestedWrapper())
3672 return emitOpError() << "must wrap an 'omp.loop_nest' directly";
3673
3674 return success();
3675}
3676
3677//===----------------------------------------------------------------------===//
3678// Distribute construct [2.9.4.1]
3679//===----------------------------------------------------------------------===//
3680
3681void DistributeOp::build(OpBuilder &builder, OperationState &state,
3682 const DistributeOperands &clauses) {
3683 DistributeOp::build(
3684 builder, state, clauses.allocateVars, clauses.allocatorVars,
3685 makeDenseI64ArrayAttr(builder.getContext(), clauses.allocateAlignments),
3687 clauses.allocatePrivateIndices),
3688 clauses.distScheduleStatic, clauses.distScheduleChunkSize, clauses.order,
3689 clauses.orderMod, clauses.privateVars,
3690 makeArrayAttr(builder.getContext(), clauses.privateSyms),
3691 clauses.privateNeedsBarrier);
3692}
3693
3694LogicalResult DistributeOp::verify() {
3695 if (this->getDistScheduleChunkSize() && !this->getDistScheduleStatic())
3696 return emitOpError() << "chunk size set without "
3697 "dist_schedule_static being present";
3698
3700 getOperation(), getAllocateVars(), getAllocatorVars(),
3701 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3702 getPrivateVars(), getPrivateSymsAttr())))
3703 return failure();
3704
3705 if (failed(verifyPrivateVarList(*this)))
3706 return failure();
3707
3708 return success();
3709}
3710
3711LogicalResult DistributeOp::verifyRegions() {
3712 if (LoopWrapperInterface nested = getNestedWrapper()) {
3713 if (!isComposite())
3714 return emitError()
3715 << "'omp.composite' attribute missing from composite wrapper";
3716 // Check for the allowed leaf constructs that may appear in a composite
3717 // construct directly after DISTRIBUTE.
3718 if (isa<WsloopOp>(nested)) {
3719 Operation *parentOp = (*this)->getParentOp();
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 "
3724 "parent";
3725 }
3726 } else if (!isa<SimdOp>(nested))
3727 return emitError() << "only supported nested wrappers are 'omp.simd' and "
3728 "'omp.wsloop'";
3729 } else if (isComposite()) {
3730 return emitError()
3731 << "'omp.composite' attribute present in non-composite wrapper";
3732 }
3733
3734 return success();
3735}
3736
3737//===----------------------------------------------------------------------===//
3738// DeclareMapperOp / DeclareMapperInfoOp
3739//===----------------------------------------------------------------------===//
3740
3741void DeclareMapperInfoOp::build(OpBuilder &builder, OperationState &state,
3742 const DeclareMapperInfoOperands &clauses) {
3743 DeclareMapperInfoOp::build(builder, state, clauses.mapVars,
3744 clauses.mapIterated);
3745}
3746
3747LogicalResult DeclareMapperInfoOp::verify() {
3748 return verifyMapClause(*this, getMapVars(), getMapIterated());
3749}
3750
3751LogicalResult DeclareMapperOp::verifyRegions() {
3752 if (!llvm::isa_and_present<DeclareMapperInfoOp>(
3753 getRegion().getBlocks().front().getTerminator()))
3754 return emitOpError() << "expected terminator to be a DeclareMapperInfoOp";
3755
3756 return success();
3757}
3758
3759//===----------------------------------------------------------------------===//
3760// DeclareReductionOp
3761//===----------------------------------------------------------------------===//
3762
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";
3770 }
3771 }
3772
3773 if (getInitializerRegion().empty())
3774 return emitOpError() << "expects non-empty initializer region";
3775 Block &initializerEntryBlock = getInitializerRegion().front();
3776
3777 if (initializerEntryBlock.getNumArguments() == 1) {
3778 if (!getAllocRegion().empty())
3779 return emitOpError() << "expects two arguments to the initializer region "
3780 "when an allocation region is used";
3781 } else if (initializerEntryBlock.getNumArguments() == 2) {
3782 if (getAllocRegion().empty())
3783 return emitOpError() << "expects one argument to the initializer region "
3784 "when no allocation region is used";
3785 } else {
3786 return emitOpError()
3787 << "expects one or two arguments to the initializer region";
3788 }
3789
3790 for (mlir::Value arg : initializerEntryBlock.getArguments())
3791 if (arg.getType() != getType())
3792 return emitOpError() << "expects initializer region argument to match "
3793 "the reduction type";
3794
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";
3800 }
3801
3802 if (getReductionRegion().empty())
3803 return emitOpError() << "expects non-empty reduction region";
3804 Block &reductionEntryBlock = getReductionRegion().front();
3805 if (reductionEntryBlock.getNumArguments() != 2 ||
3806 reductionEntryBlock.getArgumentTypes()[0] !=
3807 reductionEntryBlock.getArgumentTypes()[1] ||
3808 reductionEntryBlock.getArgumentTypes()[0] != getType())
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";
3816 }
3817
3818 if (!getAtomicReductionRegion().empty()) {
3819 Block &atomicReductionEntryBlock = getAtomicReductionRegion().front();
3820 if (atomicReductionEntryBlock.getNumArguments() != 2 ||
3821 atomicReductionEntryBlock.getArgumentTypes()[0] !=
3822 atomicReductionEntryBlock.getArgumentTypes()[1])
3823 return emitOpError() << "expects atomic reduction region with two "
3824 "arguments of the same type";
3825 auto ptrType = llvm::dyn_cast<PointerLikeType>(
3826 atomicReductionEntryBlock.getArgumentTypes()[0]);
3827 if (!ptrType ||
3828 (ptrType.getElementType() && ptrType.getElementType() != getType()))
3829 return emitOpError() << "expects atomic reduction region arguments to "
3830 "be accumulators containing the reduction type";
3831 }
3832
3833 if (getCleanupRegion().empty())
3834 return success();
3835 Block &cleanupEntryBlock = getCleanupRegion().front();
3836 if (cleanupEntryBlock.getNumArguments() != 1 ||
3837 cleanupEntryBlock.getArgument(0).getType() != getType())
3838 return emitOpError() << "expects cleanup region with one argument "
3839 "of the reduction type";
3840
3841 return success();
3842}
3843
3844//===----------------------------------------------------------------------===//
3845// TaskOp
3846//===----------------------------------------------------------------------===//
3847
3848void TaskOp::build(OpBuilder &builder, OperationState &state,
3849 const TaskOperands &clauses) {
3850 MLIRContext *ctx = builder.getContext();
3851 TaskOp::build(builder, state, clauses.iterated, clauses.affinityVars,
3852 clauses.allocateVars, clauses.allocatorVars,
3853 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3854 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3855 makeArrayAttr(ctx, clauses.dependKinds), clauses.dependVars,
3856 makeArrayAttr(ctx, clauses.dependIteratedKinds),
3857 clauses.dependIterated, clauses.final, clauses.ifExpr,
3858 clauses.inReductionVars,
3859 makeDenseBoolArrayAttr(ctx, clauses.inReductionByref),
3860 makeArrayAttr(ctx, clauses.inReductionSyms), clauses.mergeable,
3861 clauses.priority, /*private_vars=*/clauses.privateVars,
3862 /*private_syms=*/makeArrayAttr(ctx, clauses.privateSyms),
3863 clauses.privateNeedsBarrier, clauses.threadset, clauses.untied,
3864 clauses.eventHandle);
3865}
3866
3867LogicalResult TaskOp::verify() {
3869 getOperation(), getAllocateVars(), getAllocatorVars(),
3870 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3871 getPrivateVars(), getPrivateSymsAttr())))
3872 return failure();
3873
3874 LogicalResult verifyDependVars =
3875 verifyDependVarList(*this, getDependKinds(), getDependVars(),
3876 getDependIteratedKinds(), getDependIterated());
3877 if (failed(verifyDependVars))
3878 return verifyDependVars;
3879
3880 if (failed(verifyPrivateVarList(*this)))
3881 return failure();
3882
3883 return verifyReductionVarList(*this, getInReductionSyms(),
3884 getInReductionVars(), getInReductionByref());
3885}
3886
3887//===----------------------------------------------------------------------===//
3888// TaskgroupOp
3889//===----------------------------------------------------------------------===//
3890
3891void TaskgroupOp::build(OpBuilder &builder, OperationState &state,
3892 const TaskgroupOperands &clauses) {
3893 MLIRContext *ctx = builder.getContext();
3894 TaskgroupOp::build(builder, state, clauses.allocateVars,
3895 clauses.allocatorVars,
3896 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3897 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices),
3898 clauses.taskReductionVars,
3899 makeDenseBoolArrayAttr(ctx, clauses.taskReductionByref),
3900 makeArrayAttr(ctx, clauses.taskReductionSyms));
3901}
3902
3903LogicalResult TaskgroupOp::verify() {
3905 getOperation(), getAllocateVars(), getAllocatorVars(),
3906 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr())))
3907 return failure();
3908
3909 return verifyReductionVarList(*this, getTaskReductionSyms(),
3910 getTaskReductionVars(),
3911 getTaskReductionByref());
3912}
3913
3914//===----------------------------------------------------------------------===//
3915// TaskloopContextOp
3916//===----------------------------------------------------------------------===//
3917
3918void TaskloopContextOp::build(OpBuilder &builder, OperationState &state,
3919 const TaskloopContextOperands &clauses) {
3920 MLIRContext *ctx = builder.getContext();
3921 TaskloopContextOp::build(
3922 builder, state, clauses.allocateVars, clauses.allocatorVars,
3923 makeDenseI64ArrayAttr(ctx, clauses.allocateAlignments),
3924 makeDenseI64ArrayAttr(ctx, clauses.allocatePrivateIndices), clauses.final,
3925 clauses.grainsizeMod, clauses.grainsize, clauses.ifExpr,
3926 clauses.inReductionVars,
3927 makeDenseBoolArrayAttr(ctx, clauses.inReductionByref),
3928 makeArrayAttr(ctx, clauses.inReductionSyms), clauses.mergeable,
3929 clauses.nogroup, clauses.numTasksMod, clauses.numTasks, clauses.priority,
3930 /*private_vars=*/clauses.privateVars,
3931 /*private_syms=*/makeArrayAttr(ctx, clauses.privateSyms),
3932 clauses.privateNeedsBarrier, clauses.reductionMod, clauses.reductionVars,
3933 makeDenseBoolArrayAttr(ctx, clauses.reductionByref),
3934 makeArrayAttr(ctx, clauses.reductionSyms), clauses.threadset,
3935 clauses.untied);
3936 state.addAttribute("omp.combined", UnitAttr::get(ctx));
3937}
3938
3939TaskloopWrapperOp TaskloopContextOp::getLoopOp() {
3940 return cast<TaskloopWrapperOp>(
3941 *llvm::find_if(getRegion().front(), [](mlir::Operation &op) {
3942 return isa<TaskloopWrapperOp>(op);
3943 }));
3944}
3945
3946LogicalResult TaskloopContextOp::verify() {
3947 if (failed(verifyPrivateVarList(*this)))
3948 return failure();
3950 getOperation(), getAllocateVars(), getAllocatorVars(),
3951 getAllocateAlignmentsAttr(), getAllocatePrivateIndicesAttr(),
3952 getPrivateVars(), getPrivateSymsAttr())))
3953 return failure();
3954
3955 if (failed(verifyReductionVarList(*this, getReductionSyms(),
3956 getReductionVars(), getReductionByref())) ||
3957 failed(verifyReductionVarList(*this, getInReductionSyms(),
3958 getInReductionVars(),
3959 getInReductionByref())))
3960 return failure();
3961
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");
3969 }
3970
3971 if (getGrainsize() && getNumTasks()) {
3972 return emitError(
3973 "the grainsize clause and num_tasks clause are mutually exclusive and "
3974 "may not appear on the same taskloop directive");
3975 }
3976
3977 // Without this restriction, any compound construct including `taskloop` would
3978 // fail to correctly identify the whole chain of operations (see
3979 // ComposableOpInterface::findCapturedOp()), as well as failing to do so even
3980 // for standalone `taskloop` constructs.
3981 if (!isCombined())
3982 return emitOpError("must always contain the 'omp.combined' attribute");
3983
3984 return success();
3985}
3986
3987LogicalResult TaskloopContextOp::verifyRegions() {
3988 Region &region = getRegion();
3989 auto loopWrapperIt = llvm::find_if(region.front(), [](mlir::Operation &op) {
3990 return isa<TaskloopWrapperOp>(op);
3991 });
3992 if (loopWrapperIt == region.front().end())
3993 return emitOpError()
3994 << "expected a TaskloopWrapperOp directly nested in the region";
3995
3996 auto loopWrapperOp = cast<TaskloopWrapperOp>(*loopWrapperIt);
3997 auto loopNestOp = dyn_cast<LoopNestOp>(loopWrapperOp.getWrappedLoop());
3998 // This will fail the verifier for TaskloopWrapperOp and print an error
3999 // message there.
4000 if (!loopNestOp)
4001 return failure();
4002
4003 std::function<bool(Value)> isValidBoundValue = [&](Value value) -> bool {
4004 Region *valueRegion = value.getParentRegion();
4005 // A loop bound value defined outside of the taskloop context region is
4006 // valid. A region is considered an ancestor of itself.
4007 if (!region.isAncestor(valueRegion))
4008 return true;
4009
4010 Operation *defOp = value.getDefiningOp();
4011 if (!defOp || defOp->getNumRegions() != 0 || !isPure(defOp))
4012 return false;
4013
4014 return llvm::all_of(defOp->getOperands(), isValidBoundValue);
4015 };
4016 auto hasUnsupportedTaskloopLocalBound = [&](OperandRange range) -> bool {
4017 return llvm::any_of(range,
4018 [&](Value value) { return !isValidBoundValue(value); });
4019 };
4020
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";
4028 }
4029
4030 return success();
4031}
4032
4033//===----------------------------------------------------------------------===//
4034// TaskloopWrapperOp
4035//===----------------------------------------------------------------------===//
4036
4037void TaskloopWrapperOp::build(OpBuilder &builder, OperationState &state,
4038 const TaskloopWrapperOperands &clauses) {
4039 TaskloopWrapperOp::build(builder, state);
4040}
4041
4042TaskloopContextOp TaskloopWrapperOp::getTaskloopContext() {
4043 return dyn_cast<TaskloopContextOp>(getOperation()->getParentOp());
4044}
4045
4046LogicalResult TaskloopWrapperOp::verify() {
4047 TaskloopContextOp context = getTaskloopContext();
4048 if (!context)
4049 return emitOpError() << "expected to be nested in a taskloop context op";
4050 return success();
4051}
4052
4053LogicalResult TaskloopWrapperOp::verifyRegions() {
4054 if (LoopWrapperInterface nested = getNestedWrapper()) {
4055 if (!isComposite())
4056 return emitError()
4057 << "'omp.composite' attribute missing from composite wrapper";
4058
4059 // Check for the allowed leaf constructs that may appear in a composite
4060 // construct directly after TASKLOOP.
4061 if (!isa<SimdOp>(nested))
4062 return emitError() << "only supported nested wrapper is 'omp.simd'";
4063 } else if (isComposite()) {
4064 return emitError()
4065 << "'omp.composite' attribute present in non-composite wrapper";
4066 }
4067
4068 return success();
4069}
4070
4071//===----------------------------------------------------------------------===//
4072// LoopNestOp
4073//===----------------------------------------------------------------------===//
4074
4075ParseResult LoopNestOp::parse(OpAsmParser &parser, OperationState &result) {
4076 // Parse an opening `(` followed by induction variables followed by `)`
4079 Type loopVarType;
4081 parser.parseColonType(loopVarType) ||
4082 // Parse loop bounds.
4083 parser.parseEqual() ||
4084 parser.parseOperandList(lbs, ivs.size(), OpAsmParser::Delimiter::Paren) ||
4085 parser.parseKeyword("to") ||
4086 parser.parseOperandList(ubs, ivs.size(), OpAsmParser::Delimiter::Paren))
4087 return failure();
4088
4089 for (auto &iv : ivs)
4090 iv.type = loopVarType;
4091
4092 auto *ctx = parser.getBuilder().getContext();
4093 // Parse "inclusive" flag.
4094 if (succeeded(parser.parseOptionalKeyword("inclusive")))
4095 result.addAttribute("loop_inclusive", UnitAttr::get(ctx));
4096
4097 // Parse step values.
4099 if (parser.parseKeyword("step") ||
4100 parser.parseOperandList(steps, ivs.size(), OpAsmParser::Delimiter::Paren))
4101 return failure();
4102
4103 // Parse collapse
4104 int64_t value = 0;
4105 if (!parser.parseOptionalKeyword("collapse") &&
4106 (parser.parseLParen() || parser.parseInteger(value) ||
4107 parser.parseRParen()))
4108 return failure();
4109 if (value > 1)
4110 result.addAttribute(
4111 "collapse_num_loops",
4112 IntegerAttr::get(parser.getBuilder().getI64Type(), value));
4113
4114 // Parse tiles
4116 auto parseTiles = [&]() -> ParseResult {
4117 int64_t tile;
4118 if (parser.parseInteger(tile))
4119 return failure();
4120 tiles.push_back(tile);
4121 return success();
4122 };
4123
4124 if (!parser.parseOptionalKeyword("tiles") &&
4125 (parser.parseLParen() || parser.parseCommaSeparatedList(parseTiles) ||
4126 parser.parseRParen()))
4127 return failure();
4128
4129 if (tiles.size() > 0)
4130 result.addAttribute("tile_sizes", DenseI64ArrayAttr::get(ctx, tiles));
4131
4132 // Parse the body.
4133 Region *region = result.addRegion();
4134 if (parser.parseRegion(*region, ivs))
4135 return failure();
4136
4137 // Resolve operands.
4138 if (parser.resolveOperands(lbs, loopVarType, result.operands) ||
4139 parser.resolveOperands(ubs, loopVarType, result.operands) ||
4140 parser.resolveOperands(steps, loopVarType, result.operands))
4141 return failure();
4142
4143 // Parse the optional attribute list.
4144 return parser.parseOptionalAttrDict(result.attributes);
4145}
4146
4147void LoopNestOp::print(OpAsmPrinter &p) {
4148 Region &region = getRegion();
4149 auto args = region.getArguments();
4150 p << " (" << args << ") : " << args[0].getType() << " = ("
4151 << getLoopLowerBounds() << ") to (" << getLoopUpperBounds() << ") ";
4152 if (getLoopInclusive())
4153 p << "inclusive ";
4154 p << "step (" << getLoopSteps() << ") ";
4155 if (int64_t numCollapse = getCollapseNumLoops())
4156 if (numCollapse > 1)
4157 p << "collapse(" << numCollapse << ") ";
4158
4159 if (const auto tiles = getTileSizes())
4160 p << "tiles(" << tiles.value() << ") ";
4161
4162 p.printRegion(region, /*printEntryBlockArgs=*/false);
4163}
4164
4165void LoopNestOp::build(OpBuilder &builder, OperationState &state,
4166 const LoopNestOperands &clauses) {
4167 MLIRContext *ctx = builder.getContext();
4168 LoopNestOp::build(builder, state, clauses.collapseNumLoops,
4169 clauses.loopLowerBounds, clauses.loopUpperBounds,
4170 clauses.loopSteps, clauses.loopInclusive,
4171 makeDenseI64ArrayAttr(ctx, clauses.tileSizes));
4172}
4173
4174LogicalResult LoopNestOp::verify() {
4175 if (getLoopLowerBounds().empty())
4176 return emitOpError() << "must represent at least one loop";
4177
4178 if (getLoopLowerBounds().size() != getIVs().size())
4179 return emitOpError() << "number of range arguments and IVs do not match";
4180
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";
4185 }
4186
4187 uint64_t numIVs = getIVs().size();
4188
4189 if (const auto &numCollapse = getCollapseNumLoops())
4190 if (numCollapse > numIVs)
4191 return emitOpError()
4192 << "collapse value is larger than the number of loops";
4193
4194 if (const auto &tiles = getTileSizes())
4195 if (tiles.value().size() > numIVs)
4196 return emitOpError() << "too few canonical loops for tile dimensions";
4197
4198 if (!llvm::dyn_cast_if_present<LoopWrapperInterface>((*this)->getParentOp()))
4199 return emitOpError() << "expects parent op to be a loop wrapper";
4200
4201 return success();
4202}
4203
4204void LoopNestOp::gatherWrappers(
4206 Operation *parent = (*this)->getParentOp();
4207 while (auto wrapper =
4208 llvm::dyn_cast_if_present<LoopWrapperInterface>(parent)) {
4209 wrappers.push_back(wrapper);
4210 parent = parent->getParentOp();
4211 }
4212}
4213
4214//===----------------------------------------------------------------------===//
4215// OpenMP canonical loop handling
4216//===----------------------------------------------------------------------===//
4217
4218std::tuple<NewCliOp, OpOperand *, OpOperand *>
4219mlir::omp ::decodeCli(Value cli) {
4220
4221 // Defining a CLI for a generated loop is optional; if there is none then
4222 // there is no followup-tranformation
4223 if (!cli)
4224 return {{}, nullptr, nullptr};
4225
4226 assert(cli.getType() == CanonicalLoopInfoType::get(cli.getContext()) &&
4227 "Unexpected type of cli");
4228
4229 NewCliOp create = cast<NewCliOp>(cli.getDefiningOp());
4230 OpOperand *gen = nullptr;
4231 OpOperand *cons = nullptr;
4232 for (OpOperand &use : cli.getUses()) {
4233 auto op = cast<LoopTransformationInterface>(use.getOwner());
4234
4235 unsigned opnum = use.getOperandNumber();
4236 if (op.isGeneratee(opnum)) {
4237 assert(!gen && "Each CLI may have at most one def");
4238 gen = &use;
4239 } else if (op.isApplyee(opnum)) {
4240 assert(!cons && "Each CLI may have at most one consumer");
4241 cons = &use;
4242 } else {
4243 llvm_unreachable("Unexpected operand for a CLI");
4244 }
4245 }
4246
4247 return {create, gen, cons};
4248}
4249
4250ClauseProcBindKind
4251mlir::omp::convertProcBindKind(llvm::omp::ProcBindKind kind) {
4252 switch (kind) {
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:
4263 break;
4264 }
4265 llvm_unreachable("unexpected proc-bind kind");
4266}
4267
4268void NewCliOp::build(::mlir::OpBuilder &odsBuilder,
4269 ::mlir::OperationState &odsState) {
4270 odsState.addTypes(CanonicalLoopInfoType::get(odsBuilder.getContext()));
4271}
4272
4273void NewCliOp::getAsmResultNames(OpAsmSetValueNameFn setNameFn) {
4274 Value result = getResult();
4275 auto [newCli, gen, cons] = decodeCli(result);
4276
4277 // Structured binding `gen` cannot be captured in lambdas before C++20
4278 OpOperand *generator = gen;
4279
4280 // Derive the CLI variable name from its generator:
4281 // * "canonloop" for omp.canonical_loop
4282 // * custom name for loop transformation generatees
4283 // * "cli" as fallback if no generator
4284 // * "_r<idx>" suffix for nested loops, where <idx> is the sequential order
4285 // at that level
4286 // * "_s<idx>" suffix for operations with multiple regions, where <idx> is
4287 // the index of that region
4288 std::string cliName{"cli"};
4289 if (gen) {
4290 cliName =
4292 .Case([&](CanonicalLoopOp op) {
4293 return generateLoopNestingName("canonloop", op);
4294 })
4295 .Case([&](UnrollHeuristicOp op) -> std::string {
4296 llvm_unreachable("heuristic unrolling does not generate a loop");
4297 })
4298 .Case([&](FuseOp op) -> std::string {
4299 unsigned opnum = generator->getOperandNumber();
4300 // The position of the first loop to be fused is the same position
4301 // as the resulting fused loop
4302 if (op.getFirst().has_value() && opnum != op.getFirst().value())
4303 return "canonloop_fuse";
4304 else
4305 return "fused";
4306 })
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();
4314 // In the OpenMP apply and looprange clauses, indices are 1-based
4315 if (firstGrid <= opnum && opnum < firstIntratile) {
4316 unsigned gridnum = opnum - firstGrid + 1;
4317 return ("grid" + Twine(gridnum)).str();
4318 }
4319 if (firstIntratile <= opnum && opnum < end) {
4320 unsigned intratilenum = opnum - firstIntratile + 1;
4321 return ("intratile" + Twine(intratilenum)).str();
4322 }
4323 llvm_unreachable("Unexpected generatee argument");
4324 })
4325 .DefaultUnreachable("TODO: Custom name for this operation");
4326 }
4327
4328 setNameFn(result, cliName);
4329}
4330
4331LogicalResult NewCliOp::verify() {
4332 Value cli = getResult();
4333
4334 assert(cli.getType() == CanonicalLoopInfoType::get(cli.getContext()) &&
4335 "Unexpected type of cli");
4336
4337 // Check that the CLI is used in at most generator and one consumer
4338 OpOperand *gen = nullptr;
4339 OpOperand *cons = nullptr;
4340 for (mlir::OpOperand &use : cli.getUses()) {
4341 auto op = cast<mlir::omp::LoopTransformationInterface>(use.getOwner());
4342
4343 unsigned opnum = use.getOperandNumber();
4344 if (op.isGeneratee(opnum)) {
4345 if (gen) {
4346 InFlightDiagnostic error =
4347 emitOpError("CLI must have at most one generator");
4348 error.attachNote(gen->getOwner()->getLoc())
4349 .append("first generator here:");
4350 error.attachNote(use.getOwner()->getLoc())
4351 .append("second generator here:");
4352 return error;
4353 }
4354
4355 gen = &use;
4356 } else if (op.isApplyee(opnum)) {
4357 if (cons) {
4358 InFlightDiagnostic error =
4359 emitOpError("CLI must have at most one consumer");
4360 error.attachNote(cons->getOwner()->getLoc())
4361 .append("first consumer here:")
4362 .appendOp(*cons->getOwner(),
4363 OpPrintingFlags().printGenericOpForm());
4364 error.attachNote(use.getOwner()->getLoc())
4365 .append("second consumer here:")
4366 .appendOp(*use.getOwner(), OpPrintingFlags().printGenericOpForm());
4367 return error;
4368 }
4369
4370 cons = &use;
4371 } else {
4372 llvm_unreachable("Unexpected operand for a CLI");
4373 }
4374 }
4375
4376 // If the CLI is source of a transformation, it must have a generator
4377 if (cons && !gen) {
4378 InFlightDiagnostic error = emitOpError("CLI has no generator");
4379 error.attachNote(cons->getOwner()->getLoc())
4380 .append("see consumer here: ")
4381 .appendOp(*cons->getOwner(), OpPrintingFlags().printGenericOpForm());
4382 return error;
4383 }
4384
4385 return success();
4386}
4387
4388void CanonicalLoopOp::build(OpBuilder &odsBuilder, OperationState &odsState,
4389 Value tripCount) {
4390 odsState.addOperands(tripCount);
4391 odsState.addOperands(Value());
4392 (void)odsState.addRegion();
4393}
4394
4395void CanonicalLoopOp::build(OpBuilder &odsBuilder, OperationState &odsState,
4396 Value tripCount, ::mlir::Value cli) {
4397 odsState.addOperands(tripCount);
4398 odsState.addOperands(cli);
4399 (void)odsState.addRegion();
4400}
4401
4402void CanonicalLoopOp::getAsmBlockNames(OpAsmSetBlockNameFn setNameFn) {
4403 setNameFn(&getRegion().front(), "body_entry");
4404}
4405
4406void CanonicalLoopOp::getAsmBlockArgumentNames(Region &region,
4407 OpAsmSetValueNameFn setNameFn) {
4408 std::string ivName = generateLoopNestingName("iv", *this);
4409 setNameFn(region.getArgument(0), ivName);
4410}
4411
4412void CanonicalLoopOp::print(OpAsmPrinter &p) {
4413 if (getCli())
4414 p << '(' << getCli() << ')';
4415 p << ' ' << getInductionVar() << " : " << getInductionVar().getType()
4416 << " in range(" << getTripCount() << ") ";
4417
4418 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
4419 /*printBlockTerminators=*/true);
4420
4421 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
4422}
4423
4424mlir::ParseResult CanonicalLoopOp::parse(::mlir::OpAsmParser &parser,
4426 CanonicalLoopInfoType cliType =
4427 CanonicalLoopInfoType::get(parser.getContext());
4428
4429 // Parse (optional) omp.cli identifier
4431 SmallVector<mlir::Value, 1> cliOperand;
4432 if (!parser.parseOptionalLParen()) {
4433 if (parser.parseOperand(cli) ||
4434 parser.resolveOperand(cli, cliType, cliOperand) || parser.parseRParen())
4435 return failure();
4436 }
4437
4438 // We derive the type of tripCount from inductionVariable. MLIR requires the
4439 // type of tripCount to be known when calling resolveOperand so we have parse
4440 // the type before processing the inductionVariable.
4441 OpAsmParser::Argument inductionVariable;
4443 if (parser.parseArgument(inductionVariable, /*allowType*/ true) ||
4444 parser.parseKeyword("in") || parser.parseKeyword("range") ||
4445 parser.parseLParen() || parser.parseOperand(tripcount) ||
4446 parser.parseRParen() ||
4447 parser.resolveOperand(tripcount, inductionVariable.type, result.operands))
4448 return failure();
4449
4450 // Parse the loop body.
4451 Region *region = result.addRegion();
4452 if (parser.parseRegion(*region, {inductionVariable}))
4453 return failure();
4454
4455 // We parsed the cli operand forst, but because it is optional, it must be
4456 // last in the operand list.
4457 result.operands.append(cliOperand);
4458
4459 // Parse the optional attribute list.
4460 if (parser.parseOptionalAttrDict(result.attributes))
4461 return failure();
4462
4463 return mlir::success();
4464}
4465
4466LogicalResult CanonicalLoopOp::verify() {
4467 // The region's entry must accept the induction variable
4468 // It can also be empty if just created
4469 if (!getRegion().empty()) {
4470 Region &region = getRegion();
4471 if (region.getNumArguments() != 1)
4472 return emitOpError(
4473 "Canonical loop region must have exactly one argument");
4474
4475 if (getInductionVar().getType() != getTripCount().getType())
4476 return emitOpError(
4477 "Region argument must be the same type as the trip count");
4478 }
4479
4480 return success();
4481}
4482
4483Value CanonicalLoopOp::getInductionVar() { return getRegion().getArgument(0); }
4484
4485std::pair<unsigned, unsigned>
4486CanonicalLoopOp::getApplyeesODSOperandIndexAndLength() {
4487 // No applyees
4488 return {0, 0};
4489}
4490
4491std::pair<unsigned, unsigned>
4492CanonicalLoopOp::getGenerateesODSOperandIndexAndLength() {
4493 return getODSOperandIndexAndLength(odsIndex_cli);
4494}
4495
4496//===----------------------------------------------------------------------===//
4497// UnrollHeuristicOp
4498//===----------------------------------------------------------------------===//
4499
4500void UnrollHeuristicOp::build(::mlir::OpBuilder &odsBuilder,
4501 ::mlir::OperationState &odsState,
4502 ::mlir::Value cli) {
4503 odsState.addOperands(cli);
4504}
4505
4506void UnrollHeuristicOp::print(OpAsmPrinter &p) {
4507 p << '(' << getApplyee() << ')';
4508
4509 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
4510}
4511
4512mlir::ParseResult UnrollHeuristicOp::parse(::mlir::OpAsmParser &parser,
4514 auto cliType = CanonicalLoopInfoType::get(parser.getContext());
4515
4516 if (parser.parseLParen())
4517 return failure();
4518
4520 if (parser.parseOperand(applyee) ||
4521 parser.resolveOperand(applyee, cliType, result.operands))
4522 return failure();
4523
4524 if (parser.parseRParen())
4525 return failure();
4526
4527 // Optional output loop (full unrolling has none)
4528 if (!parser.parseOptionalArrow()) {
4529 if (parser.parseLParen() || parser.parseRParen())
4530 return failure();
4531 }
4532
4533 // Parse the optional attribute list.
4534 if (parser.parseOptionalAttrDict(result.attributes))
4535 return failure();
4536
4537 return mlir::success();
4538}
4539
4540std::pair<unsigned, unsigned>
4541UnrollHeuristicOp ::getApplyeesODSOperandIndexAndLength() {
4542 return getODSOperandIndexAndLength(odsIndex_applyee);
4543}
4544
4545std::pair<unsigned, unsigned>
4546UnrollHeuristicOp::getGenerateesODSOperandIndexAndLength() {
4547 return {0, 0};
4548}
4549
4550//===----------------------------------------------------------------------===//
4551// UnrollFullOp
4552//===----------------------------------------------------------------------===//
4553
4554void UnrollFullOp::build(::mlir::OpBuilder &odsBuilder,
4555 ::mlir::OperationState &odsState, ::mlir::Value cli) {
4556 odsState.addOperands(cli);
4557}
4558
4559void UnrollFullOp::print(OpAsmPrinter &p) {
4560 p << '(' << getApplyee() << ')';
4561
4562 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
4563}
4564
4565mlir::ParseResult UnrollFullOp::parse(::mlir::OpAsmParser &parser,
4567 auto cliType = CanonicalLoopInfoType::get(parser.getContext());
4568
4569 if (parser.parseLParen())
4570 return failure();
4571
4573 if (parser.parseOperand(applyee) ||
4574 parser.resolveOperand(applyee, cliType, result.operands))
4575 return failure();
4576
4577 if (parser.parseRParen())
4578 return failure();
4579
4580 // Optional output loop; full unrolling has none.
4581 if (!parser.parseOptionalArrow()) {
4582 if (parser.parseLParen() || parser.parseRParen())
4583 return failure();
4584 }
4585
4586 // Parse the optional attribute list.
4587 if (parser.parseOptionalAttrDict(result.attributes))
4588 return failure();
4589
4590 return mlir::success();
4591}
4592
4593std::pair<unsigned, unsigned>
4594UnrollFullOp::getApplyeesODSOperandIndexAndLength() {
4595 return getODSOperandIndexAndLength(odsIndex_applyee);
4596}
4597
4598std::pair<unsigned, unsigned>
4599UnrollFullOp::getGenerateesODSOperandIndexAndLength() {
4600 return {0, 0};
4601}
4602
4603LogicalResult UnrollFullOp::verify() {
4604 auto [create, gen, cons] = decodeCli(getApplyee());
4605 if (!gen)
4606 return emitOpError() << "applyee CLI has no generator";
4607
4608 // Full unrolling leaves no loop, so the trip count must be constant. Only
4609 // omp.canonical_loop states one.
4610 if (auto loop = dyn_cast<CanonicalLoopOp>(gen->getOwner())) {
4611 if (!matchPattern(loop.getTripCount(), m_Constant()))
4612 return emitOpError() << "applyee loop must have a constant trip count";
4613 }
4614
4615 return success();
4616}
4617
4618//===----------------------------------------------------------------------===//
4619// UnrollPartialOp
4620//===----------------------------------------------------------------------===//
4621
4622void UnrollPartialOp::build(::mlir::OpBuilder &odsBuilder,
4624 uint64_t unrollFactor) {
4625 odsState.addOperands(cli);
4626 Properties &props = odsState.getOrAddProperties<Properties>();
4627 props.unroll_factor = odsBuilder.getI64IntegerAttr(unrollFactor);
4628}
4629
4630void UnrollPartialOp::print(OpAsmPrinter &p) {
4631 p << '(' << getApplyee() << ')';
4632
4633 SmallVector<NamedAttribute> attrs((*this)->getDiscardableAttrs());
4634 attrs.emplace_back(getUnrollFactorAttrName(), getUnrollFactorAttr());
4635 llvm::sort(attrs);
4636 p.printOptionalAttrDict(attrs);
4637}
4638
4639mlir::ParseResult UnrollPartialOp::parse(::mlir::OpAsmParser &parser,
4641 auto cliType = CanonicalLoopInfoType::get(parser.getContext());
4642
4643 if (parser.parseLParen())
4644 return failure();
4645
4647 if (parser.parseOperand(applyee) ||
4648 parser.resolveOperand(applyee, cliType, result.operands))
4649 return failure();
4650
4651 if (parser.parseRParen())
4652 return failure();
4653
4654 // The unroll factor is carried by the `unroll_factor` attribute.
4655 if (parser.parseOptionalAttrDict(result.attributes))
4656 return failure();
4657
4658 return mlir::success();
4659}
4660
4661std::pair<unsigned, unsigned>
4662UnrollPartialOp::getApplyeesODSOperandIndexAndLength() {
4663 return getODSOperandIndexAndLength(odsIndex_applyee);
4664}
4665
4666std::pair<unsigned, unsigned>
4667UnrollPartialOp::getGenerateesODSOperandIndexAndLength() {
4668 return {0, 0};
4669}
4670
4671//===----------------------------------------------------------------------===//
4672// TileOp
4673//===----------------------------------------------------------------------===//
4674
4675static void printLoopTransformClis(OpAsmPrinter &p, TileOp op,
4676 OperandRange generatees,
4677 OperandRange applyees) {
4678 if (!generatees.empty())
4679 p << '(' << llvm::interleaved(generatees) << ')';
4680
4681 if (!applyees.empty())
4682 p << " <- (" << llvm::interleaved(applyees) << ')';
4683}
4684
4685static ParseResult parseLoopTransformClis(
4686 OpAsmParser &parser,
4689 if (parser.parseOptionalLess()) {
4690 // Syntax 1: generatees present
4691
4692 if (parser.parseOperandList(generateesOperands,
4694 return failure();
4695
4696 if (parser.parseLess())
4697 return failure();
4698 } else {
4699 // Syntax 2: generatees omitted
4700 }
4701
4702 // Parse `<-` (`<` has already been parsed)
4703 if (parser.parseMinus())
4704 return failure();
4705
4706 if (parser.parseOperandList(applyeesOperands,
4708 return failure();
4709
4710 return success();
4711}
4712
4713/// Check properties of the loop nest consisting of the transformation's
4714/// applyees:
4715/// 1. They are nested inside each other
4716/// 2. They are perfectly nested
4717/// (no code with side-effects in-between the loops)
4718/// 3. They are rectangular
4719/// (loop bounds are invariant in respect to the outer loops)
4720///
4721/// TODO: Generalize for LoopTransformationInterface.
4722static LogicalResult checkApplyeesNesting(TileOp op) {
4723 // Collect the loops from the nest
4724 bool isOnlyCanonLoops = true;
4726 for (Value applyee : op.getApplyees()) {
4727 auto [create, gen, cons] = decodeCli(applyee);
4728
4729 if (!gen)
4730 return op.emitOpError() << "applyee CLI has no generator";
4731
4732 auto loop = dyn_cast_or_null<CanonicalLoopOp>(gen->getOwner());
4733 canonLoops.push_back(loop);
4734 if (!loop)
4735 isOnlyCanonLoops = false;
4736 }
4737
4738 // FIXME: We currently can only verify non-rectangularity and perfect nest of
4739 // omp.canonical_loop.
4740 if (!isOnlyCanonLoops)
4741 return success();
4742
4743 DenseSet<Value> parentIVs;
4744 for (auto i : llvm::seq<int>(1, canonLoops.size())) {
4745 auto parentLoop = canonLoops[i - 1];
4746 auto loop = canonLoops[i];
4747
4748 if (parentLoop.getOperation() != loop.getOperation()->getParentOp())
4749 return op.emitOpError()
4750 << "tiled loop nest must be nested within each other";
4751
4752 parentIVs.insert(parentLoop.getInductionVar());
4753
4754 // Canonical loop must be perfectly nested, i.e. the body of the parent must
4755 // only contain the omp.canonical_loop of the nested loops, and
4756 // omp.terminator
4757 bool isPerfectlyNested = [&]() {
4758 auto &parentBody = parentLoop.getRegion();
4759 if (!parentBody.hasOneBlock())
4760 return false;
4761 auto &parentBlock = parentBody.getBlocks().front();
4762
4763 auto nestedLoopIt = parentBlock.begin();
4764 if (nestedLoopIt == parentBlock.end() ||
4765 (&*nestedLoopIt != loop.getOperation()))
4766 return false;
4767
4768 auto termIt = std::next(nestedLoopIt);
4769 if (termIt == parentBlock.end() || !isa<TerminatorOp>(termIt))
4770 return false;
4771
4772 if (std::next(termIt) != parentBlock.end())
4773 return false;
4774
4775 return true;
4776 }();
4777 if (!isPerfectlyNested)
4778 return op.emitOpError() << "tiled loop nest must be perfectly nested";
4779
4780 if (parentIVs.contains(loop.getTripCount()))
4781 return op.emitOpError() << "tiled loop nest must be rectangular";
4782 }
4783
4784 // TODO: The tile sizes must be computed before the loop, but checking this
4785 // requires dominance analysis. For instance:
4786 //
4787 // %canonloop = omp.new_cli
4788 // omp.canonical_loop(%canonloop) %iv : i32 in range(%tc) {
4789 // // write to %x
4790 // omp.terminator
4791 // }
4792 // %ts = llvm.load %x
4793 // omp.tile <- (%canonloop) sizes(%ts : i32)
4794
4795 return success();
4796}
4797
4798LogicalResult TileOp::verify() {
4799 if (getApplyees().empty())
4800 return emitOpError() << "must apply to at least one loop";
4801
4802 if (getSizes().size() != getApplyees().size())
4803 return emitOpError() << "there must be one tile size for each applyee";
4804
4805 if (!getGeneratees().empty() &&
4806 2 * getSizes().size() != getGeneratees().size())
4807 return emitOpError()
4808 << "expecting two times the number of generatees than applyees";
4809
4810 return checkApplyeesNesting(*this);
4811}
4812
4813std::pair<unsigned, unsigned> TileOp ::getApplyeesODSOperandIndexAndLength() {
4814 return getODSOperandIndexAndLength(odsIndex_applyees);
4815}
4816
4817std::pair<unsigned, unsigned> TileOp::getGenerateesODSOperandIndexAndLength() {
4818 return getODSOperandIndexAndLength(odsIndex_generatees);
4819}
4820
4821//===----------------------------------------------------------------------===//
4822// FuseOp
4823//===----------------------------------------------------------------------===//
4824
4825static void printLoopTransformClis(OpAsmPrinter &p, FuseOp op,
4826 OperandRange generatees,
4827 OperandRange applyees) {
4828 if (!generatees.empty())
4829 p << '(' << llvm::interleaved(generatees) << ')';
4830
4831 if (!applyees.empty())
4832 p << " <- (" << llvm::interleaved(applyees) << ')';
4833}
4834
4835LogicalResult FuseOp::verify() {
4836 if (getApplyees().size() < 2)
4837 return emitOpError() << "must apply to at least two loops";
4838
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";
4849
4850 } else {
4851 if (!getGeneratees().empty() && getGeneratees().size() != 1)
4852 return emitOpError()
4853 << "in a complete fuse the number of generatees must be exactly 1";
4854 }
4855 for (auto &&applyee : getApplyees()) {
4856 auto [create, gen, cons] = decodeCli(applyee);
4857
4858 if (!gen)
4859 return emitOpError() << "applyee CLI has no generator";
4860 auto loop = dyn_cast_or_null<CanonicalLoopOp>(gen->getOwner());
4861 if (!loop)
4862 return emitOpError()
4863 << "currently only supports omp.canonical_loop as applyee";
4864 }
4865 return success();
4866}
4867std::pair<unsigned, unsigned> FuseOp::getApplyeesODSOperandIndexAndLength() {
4868 return getODSOperandIndexAndLength(odsIndex_applyees);
4869}
4870
4871std::pair<unsigned, unsigned> FuseOp::getGenerateesODSOperandIndexAndLength() {
4872 return getODSOperandIndexAndLength(odsIndex_generatees);
4873}
4874
4875//===----------------------------------------------------------------------===//
4876// Critical construct (2.17.1)
4877//===----------------------------------------------------------------------===//
4878
4879void CriticalDeclareOp::build(OpBuilder &builder, OperationState &state,
4880 const CriticalDeclareOperands &clauses) {
4881 CriticalDeclareOp::build(builder, state, clauses.symName,
4882 clauses.symVisibility, clauses.hint);
4883}
4884
4885LogicalResult CriticalDeclareOp::verify() {
4886 return verifySynchronizationHint(*this, getHint());
4887}
4888
4889LogicalResult CriticalOp::verify() {
4890 SymbolRefAttr currentName = getNameAttr();
4891
4892 CriticalOp parentCritical = (*this)->getParentOfType<CriticalOp>();
4893
4894 while (parentCritical) {
4895 SymbolRefAttr parentName = parentCritical.getNameAttr();
4896
4897 if (currentName == parentName) {
4898 if (currentName) {
4899 return emitOpError() << "cannot be nested inside another omp.critical "
4900 "region with the same name ("
4901 << currentName << ")";
4902 } else {
4903 return emitOpError() << "cannot be nested inside another unnamed "
4904 "omp.critical region";
4905 }
4906 }
4907
4908 parentCritical = parentCritical->getParentOfType<CriticalOp>();
4909 }
4910
4911 return success();
4912}
4913
4914LogicalResult CriticalOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
4915 if (getNameAttr()) {
4916 SymbolRefAttr symbolRef = getNameAttr();
4917 auto decl = symbolTable.lookupNearestSymbolFrom<CriticalDeclareOp>(
4918 *this, symbolRef);
4919 if (!decl) {
4920 return emitOpError() << "expected symbol reference " << symbolRef
4921 << " to point to a critical declaration";
4922 }
4923 }
4924
4925 return success();
4926}
4927
4928//===----------------------------------------------------------------------===//
4929// Spec 5.1: Error directive (2.5.4)
4930//===----------------------------------------------------------------------===//
4931
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";
4937 return success();
4938}
4939
4940//===----------------------------------------------------------------------===//
4941// Ordered construct
4942//===----------------------------------------------------------------------===//
4943
4944static LogicalResult verifyOrderedParent(Operation &op) {
4945 bool hasRegion = op.getNumRegions() > 0;
4946 auto loopOp = op.getParentOfType<LoopNestOp>();
4947 if (!loopOp) {
4948 if (hasRegion)
4949 return success();
4950
4951 // TODO: Consider if this needs to be the case only for the standalone
4952 // variant of the ordered construct.
4953 return op.emitOpError() << "must be nested inside of a loop";
4954 }
4955
4956 Operation *wrapper = loopOp->getParentOp();
4957 if (auto wsloopOp = dyn_cast<WsloopOp>(wrapper)) {
4958 IntegerAttr orderedAttr = wsloopOp.getOrderedAttr();
4959 if (!orderedAttr)
4960 return op.emitOpError() << "the enclosing worksharing-loop region must "
4961 "have an ordered clause";
4962
4963 if (hasRegion && orderedAttr.getInt() != 0)
4964 return op.emitOpError() << "the enclosing loop's ordered clause must not "
4965 "have a parameter present";
4966
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";
4973 }
4974 return success();
4975}
4976
4977void OrderedOp::build(OpBuilder &builder, OperationState &state,
4978 const OrderedOperands &clauses) {
4979 OrderedOp::build(builder, state, clauses.doacrossDependType,
4980 clauses.doacrossNumLoops, clauses.doacrossDependVars);
4981}
4982
4983LogicalResult OrderedOp::verify() {
4984 if (failed(verifyOrderedParent(**this)))
4985 return failure();
4986
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 "
4991 << "doacross loop";
4992
4993 return success();
4994}
4995
4996void OrderedRegionOp::build(OpBuilder &builder, OperationState &state,
4997 const OrderedRegionOperands &clauses) {
4998 OrderedRegionOp::build(builder, state, clauses.parLevelSimd);
4999}
5000
5001LogicalResult OrderedRegionOp::verify() { return verifyOrderedParent(**this); }
5002
5003//===----------------------------------------------------------------------===//
5004// TaskwaitOp
5005//===----------------------------------------------------------------------===//
5006
5007void TaskwaitOp::build(OpBuilder &builder, OperationState &state,
5008 const TaskwaitOperands &clauses) {
5009 // TODO Store clauses in op: depend_iterated_kinds, depend_iterated, nowait.
5010 MLIRContext *ctx = builder.getContext();
5011 TaskwaitOp::build(
5012 builder, state,
5013 /*depend_kinds=*/makeArrayAttr(ctx, clauses.dependKinds),
5014 /*depend_vars=*/clauses.dependVars,
5015 /*depend_iterated_kinds=*/makeArrayAttr(ctx, clauses.dependIteratedKinds),
5016 /*depend_iterated=*/ValueRange(clauses.dependIterated),
5017 /*nowait=*/clauses.nowait);
5018}
5019
5020//===----------------------------------------------------------------------===//
5021// Verifier for AtomicReadOp
5022//===----------------------------------------------------------------------===//
5023
5024LogicalResult AtomicReadOp::verify() {
5025 if (verifyCommon().failed())
5026 return mlir::failure();
5027
5028 int64_t version = 50;
5029 if (auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5030 if (Attribute verAttr = moduleOp->getDiscardableAttr("omp.version"))
5031 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5032
5033 if (auto mo = getMemoryOrder()) {
5034 if (*mo == ClauseMemoryOrderKind::Release) {
5035 return emitError("memory-order must not be release for atomic reads");
5036 }
5037 if (*mo == ClauseMemoryOrderKind::Acq_rel) {
5038 // acq_rel is prohibited on read only in OpenMP 5.0; allowed in 5.1+.
5039 if (version < 51)
5040 return emitError("memory-order must not be acq_rel for atomic reads");
5041 }
5042 }
5043 return verifySynchronizationHint(*this, getHint());
5044}
5045
5046//===----------------------------------------------------------------------===//
5047// Verifier for AtomicWriteOp
5048//===----------------------------------------------------------------------===//
5049
5050LogicalResult AtomicWriteOp::verify() {
5051 if (verifyCommon().failed())
5052 return mlir::failure();
5053
5054 int64_t version = 50;
5055 if (auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5056 if (Attribute verAttr = moduleOp->getDiscardableAttr("omp.version"))
5057 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5058
5059 if (auto mo = getMemoryOrder()) {
5060 if (*mo == ClauseMemoryOrderKind::Acquire) {
5061 return emitError("memory-order must not be acquire for atomic writes");
5062 }
5063 if (*mo == ClauseMemoryOrderKind::Acq_rel) {
5064 // acq_rel is prohibited on write only in OpenMP 5.0; allowed in 5.1+.
5065 if (version < 51)
5066 return emitError("memory-order must not be acq_rel for atomic writes");
5067 }
5068 }
5069 return verifySynchronizationHint(*this, getHint());
5070}
5071
5072//===----------------------------------------------------------------------===//
5073// Verifier for AtomicUpdateOp
5074//===----------------------------------------------------------------------===//
5075
5076LogicalResult AtomicUpdateOp::canonicalize(AtomicUpdateOp op,
5077 PatternRewriter &rewriter) {
5078 if (op.isNoOp()) {
5079 rewriter.eraseOp(op);
5080 return success();
5081 }
5082 if (Value writeVal = op.getWriteOpVal()) {
5083 rewriter.replaceOpWithNewOp<AtomicWriteOp>(
5084 op, op.getX(), writeVal, op.getHintAttr(), op.getMemoryOrderAttr());
5085 return success();
5086 }
5087 return failure();
5088}
5089
5090LogicalResult AtomicUpdateOp::verify() {
5091 if (verifyCommon().failed())
5092 return mlir::failure();
5093
5094 int64_t version = 50;
5095 if (auto moduleOp = getOperation()->getParentOfType<ModuleOp>())
5096 if (Attribute verAttr = moduleOp->getDiscardableAttr("omp.version"))
5097 version = llvm::cast<VersionAttr>(verAttr).getVersion();
5098
5099 if (auto mo = getMemoryOrder()) {
5100 if (*mo == ClauseMemoryOrderKind::Acq_rel ||
5101 *mo == ClauseMemoryOrderKind::Acquire) {
5102 // This restriction applies only to OpenMP 5.0; removed in 5.1.
5103 if (version < 51)
5104 return emitError(
5105 "memory-order must not be acq_rel or acquire for atomic updates");
5106 }
5107 }
5108
5109 return verifySynchronizationHint(*this, getHint());
5110}
5111
5112LogicalResult AtomicUpdateOp::verifyRegions() { return verifyRegionsCommon(); }
5113
5114//===----------------------------------------------------------------------===//
5115// Verifier for AtomicCaptureOp
5116//===----------------------------------------------------------------------===//
5117
5118AtomicReadOp AtomicCaptureOp::getAtomicReadOp() {
5119 if (auto op = dyn_cast<AtomicReadOp>(getFirstOp()))
5120 return op;
5121 return dyn_cast<AtomicReadOp>(getSecondOp());
5122}
5123
5124AtomicWriteOp AtomicCaptureOp::getAtomicWriteOp() {
5125 if (auto op = dyn_cast<AtomicWriteOp>(getFirstOp()))
5126 return op;
5127 return dyn_cast<AtomicWriteOp>(getSecondOp());
5128}
5129
5130AtomicUpdateOp AtomicCaptureOp::getAtomicUpdateOp() {
5131 if (auto op = dyn_cast<AtomicUpdateOp>(getFirstOp()))
5132 return op;
5133 return dyn_cast<AtomicUpdateOp>(getSecondOp());
5134}
5135
5136AtomicCompareOp AtomicCaptureOp::getAtomicCompareOp() {
5137 if (auto op = dyn_cast<AtomicCompareOp>(getFirstOp()))
5138 return op;
5139 return dyn_cast<AtomicCompareOp>(getSecondOp());
5140}
5141
5142LogicalResult AtomicCaptureOp::verify() {
5143 return verifySynchronizationHint(*this, getHint());
5144}
5145
5146LogicalResult AtomicCaptureOp::verifyRegions() {
5147 if (verifyRegionsCommon().failed())
5148 return mlir::failure();
5149
5150 if (getFirstOp()->getInherentAttr("hint").value_or(Attribute{}) ||
5151 getSecondOp()->getInherentAttr("hint").value_or(Attribute{}))
5152 return emitOpError(
5153 "operations inside capture region must not have hint clause");
5154
5155 if (getFirstOp()->getInherentAttr("memory_order").value_or(Attribute{}) ||
5156 getSecondOp()->getInherentAttr("memory_order").value_or(Attribute{}))
5157 return emitOpError(
5158 "operations inside capture region must not have memory_order clause");
5159 return success();
5160}
5161
5162//===----------------------------------------------------------------------===//
5163// AtomicCompareOp
5164//===----------------------------------------------------------------------===//
5165
5166LogicalResult AtomicCompareOp::verify() {
5167 if (verifyCommon().failed())
5168 return mlir::failure();
5169 // OpenMP 5.2 [15.8.3]: the fail clause argument must be one of seq_cst,
5170 // acquire or relaxed ('release' and 'acq_rel' are not valid failure
5171 // orderings and map to invalid cmpxchg failure orderings).
5172 if (auto failOrder = getFailMemoryOrder()) {
5173 if (*failOrder != ClauseMemoryOrderKind::Seq_cst &&
5174 *failOrder != ClauseMemoryOrderKind::Acquire &&
5175 *failOrder != ClauseMemoryOrderKind::Relaxed)
5176 return emitOpError(
5177 "fail_memory_order must be 'seq_cst', 'acquire' or 'relaxed'");
5178 }
5179 return verifySynchronizationHint(*this, getHint());
5180}
5181
5182LogicalResult AtomicCompareOp::verifyRegions() {
5183 if (verifyRegionsCommon().failed())
5184 return mlir::failure();
5185
5186 if (verifyOperator().failed())
5187 return mlir::failure();
5188
5189 Block &block = getRegion().front();
5190
5191 Operation *terminator = block.getTerminator();
5192 if (!terminator || !isa<YieldOp>(terminator))
5193 return emitOpError("region must be terminated with omp.yield");
5194
5195 return success();
5196}
5197
5198//===----------------------------------------------------------------------===//
5199// CancelOp
5200//===----------------------------------------------------------------------===//
5201
5202void CancelOp::build(OpBuilder &builder, OperationState &state,
5203 const CancelOperands &clauses) {
5204 CancelOp::build(builder, state, clauses.cancelDirective, clauses.ifExpr);
5205}
5206
5208 Operation *parent = thisOp->getParentOp();
5209 while (parent) {
5210 if (parent->getDialect() == thisOp->getDialect())
5211 return parent;
5212 parent = parent->getParentOp();
5213 }
5214 return nullptr;
5215}
5216
5217LogicalResult CancelOp::verify() {
5218 ClauseCancellationConstructType cct = getCancelDirective();
5219 // The next OpenMP operation in the chain of parents
5220 Operation *structuralParent = getParentInSameDialect((*this).getOperation());
5221 if (!structuralParent)
5222 return emitOpError() << "Orphaned cancel construct";
5223
5224 if ((cct == ClauseCancellationConstructType::Parallel) &&
5225 !mlir::isa<ParallelOp>(structuralParent)) {
5226 return emitOpError() << "cancel parallel must appear "
5227 << "inside a parallel region";
5228 }
5229 if (cct == ClauseCancellationConstructType::Loop) {
5230 // structural parent will be omp.loop_nest, directly nested inside
5231 // omp.wsloop
5232 auto wsloopOp = mlir::dyn_cast<WsloopOp>(structuralParent->getParentOp());
5233
5234 if (!wsloopOp) {
5235 return emitOpError()
5236 << "cancel loop must appear inside a worksharing-loop region";
5237 }
5238 if (wsloopOp.getNowaitAttr()) {
5239 return emitError() << "A worksharing construct that is canceled "
5240 << "must not have a nowait clause";
5241 }
5242 if (wsloopOp.getOrderedAttr()) {
5243 return emitError() << "A worksharing construct that is canceled "
5244 << "must not have an ordered clause";
5245 }
5246
5247 } else if (cct == ClauseCancellationConstructType::Sections) {
5248 // structural parent will be an omp.section, directly nested inside
5249 // omp.sections
5250 auto sectionsOp =
5251 mlir::dyn_cast<SectionsOp>(structuralParent->getParentOp());
5252 if (!sectionsOp) {
5253 return emitOpError() << "cancel sections must appear "
5254 << "inside a sections region";
5255 }
5256 if (sectionsOp.getNowait()) {
5257 return emitError() << "A sections construct that is canceled "
5258 << "must not have a nowait clause";
5259 }
5260 }
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";
5266 }
5267 return success();
5268}
5269
5270//===----------------------------------------------------------------------===//
5271// CancellationPointOp
5272//===----------------------------------------------------------------------===//
5273
5274void CancellationPointOp::build(OpBuilder &builder, OperationState &state,
5275 const CancellationPointOperands &clauses) {
5276 CancellationPointOp::build(builder, state, clauses.cancelDirective);
5277}
5278
5279LogicalResult CancellationPointOp::verify() {
5280 ClauseCancellationConstructType cct = getCancelDirective();
5281 // The next OpenMP operation in the chain of parents
5282 Operation *structuralParent = getParentInSameDialect((*this).getOperation());
5283 if (!structuralParent)
5284 return emitOpError() << "Orphaned cancellation point";
5285
5286 if ((cct == ClauseCancellationConstructType::Parallel) &&
5287 !mlir::isa<ParallelOp>(structuralParent)) {
5288 return emitOpError() << "cancellation point parallel must appear "
5289 << "inside a parallel region";
5290 }
5291 // Strucutal parent here will be an omp.loop_nest. Get the parent of that to
5292 // find the wsloop
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";
5297 }
5298 if ((cct == ClauseCancellationConstructType::Sections) &&
5299 !mlir::isa<omp::SectionOp>(structuralParent)) {
5300 return emitOpError() << "cancellation point sections must appear "
5301 << "inside a sections region";
5302 }
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";
5308 }
5309 return success();
5310}
5311
5312//===----------------------------------------------------------------------===//
5313// MapBoundsOp
5314//===----------------------------------------------------------------------===//
5315
5316LogicalResult MapBoundsOp::verify() {
5317 auto extent = getExtent();
5318 auto upperbound = getUpperBound();
5319 if (!extent && !upperbound)
5320 return emitError("expected extent or upperbound.");
5321 return success();
5322}
5323
5324void PrivateClauseOp::build(OpBuilder &odsBuilder, OperationState &odsState,
5325 TypeRange /*result_types*/, StringAttr symName,
5326 TypeAttr type) {
5327 PrivateClauseOp::build(
5328 odsBuilder, odsState, symName, /*sym_visibility=*/nullptr, type,
5329 DataSharingClauseTypeAttr::get(odsBuilder.getContext(),
5330 DataSharingClauseType::Private));
5331}
5332
5333LogicalResult PrivateClauseOp::verifyRegions() {
5334 Type argType = getArgType();
5335 auto verifyTerminator = [&](Operation *terminator,
5336 bool yieldsValue) -> LogicalResult {
5337 if (!terminator->getBlock()->getSuccessors().empty())
5338 return success();
5339
5340 if (!llvm::isa<YieldOp>(terminator))
5341 return mlir::emitError(terminator->getLoc())
5342 << "expected exit block terminator to be an `omp.yield` op.";
5343
5344 YieldOp yieldOp = llvm::cast<YieldOp>(terminator);
5345 TypeRange yieldedTypes = yieldOp.getResults().getTypes();
5346
5347 if (!yieldsValue) {
5348 if (yieldedTypes.empty())
5349 return success();
5350
5351 return mlir::emitError(terminator->getLoc())
5352 << "Did not expect any values to be yielded.";
5353 }
5354
5355 if (yieldedTypes.size() == 1 && yieldedTypes.front() == argType)
5356 return success();
5357
5358 auto error = mlir::emitError(yieldOp.getLoc())
5359 << "Invalid yielded value. Expected type: " << argType
5360 << ", got: ";
5361
5362 if (yieldedTypes.empty())
5363 error << "None";
5364 else
5365 error << yieldedTypes;
5366
5367 return error;
5368 };
5369
5370 auto verifyRegion = [&](Region &region, unsigned expectedNumArgs,
5371 StringRef regionName,
5372 bool yieldsValue) -> LogicalResult {
5373 assert(!region.empty());
5374
5375 if (region.getNumArguments() != expectedNumArgs)
5376 return mlir::emitError(region.getLoc())
5377 << "`" << regionName << "`: " << "expected " << expectedNumArgs
5378 << " region arguments, got: " << region.getNumArguments();
5379
5380 for (Block &block : region) {
5381 // MLIR will verify the absence of the terminator for us.
5382 if (!block.mightHaveTerminator())
5383 continue;
5384
5385 if (failed(verifyTerminator(block.getTerminator(), yieldsValue)))
5386 return failure();
5387 }
5388
5389 return success();
5390 };
5391
5392 // Ensure all of the region arguments have the same type
5393 for (Region *region : getRegions())
5394 for (Type ty : region->getArgumentTypes())
5395 if (ty != argType)
5396 return emitError() << "Region argument type mismatch: got " << ty
5397 << " expected " << argType << ".";
5398
5399 mlir::Region &initRegion = getInitRegion();
5400 if (!initRegion.empty() &&
5401 failed(verifyRegion(getInitRegion(), /*expectedNumArgs=*/2, "init",
5402 /*yieldsValue=*/true)))
5403 return failure();
5404
5405 DataSharingClauseType dsType = getDataSharingType();
5406
5407 if (dsType == DataSharingClauseType::Private && !getCopyRegion().empty())
5408 return emitError("`private` clauses do not require a `copy` region.");
5409
5410 if (dsType == DataSharingClauseType::FirstPrivate && getCopyRegion().empty())
5411 return emitError(
5412 "`firstprivate` clauses require at least a `copy` region.");
5413
5414 if (dsType == DataSharingClauseType::FirstPrivate &&
5415 failed(verifyRegion(getCopyRegion(), /*expectedNumArgs=*/2, "copy",
5416 /*yieldsValue=*/true)))
5417 return failure();
5418
5419 if (!getDeallocRegion().empty() &&
5420 failed(verifyRegion(getDeallocRegion(), /*expectedNumArgs=*/1, "dealloc",
5421 /*yieldsValue=*/false)))
5422 return failure();
5423
5424 return success();
5425}
5426
5427//===----------------------------------------------------------------------===//
5428// Spec 5.2: Masked construct (10.5)
5429//===----------------------------------------------------------------------===//
5430
5431void MaskedOp::build(OpBuilder &builder, OperationState &state,
5432 const MaskedOperands &clauses) {
5433 MaskedOp::build(builder, state, clauses.filteredThreadId);
5434}
5435
5436//===----------------------------------------------------------------------===//
5437// Spec 5.2: Dispatch construct (7.6)
5438//===----------------------------------------------------------------------===//
5439
5440void DispatchOp::build(OpBuilder &builder, OperationState &state,
5441 const DispatchOperands &clauses) {
5442 DispatchOp::build(builder, state, clauses.nocontext, clauses.novariants,
5443 clauses.nowait);
5444}
5445
5446//===----------------------------------------------------------------------===//
5447// Spec 5.2: Scan construct (5.6)
5448//===----------------------------------------------------------------------===//
5449
5450void ScanOp::build(OpBuilder &builder, OperationState &state,
5451 const ScanOperands &clauses) {
5452 ScanOp::build(builder, state, clauses.inclusiveVars, clauses.exclusiveVars);
5453}
5454
5455LogicalResult ScanOp::verify() {
5456 if (hasExclusiveVars() == hasInclusiveVars())
5457 return emitError(
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)
5463 return success();
5464 }
5465 if (SimdOp parentSimdOp = (*this)->getParentOfType<SimdOp>()) {
5466 if (parentSimdOp.getReductionModAttr() &&
5467 parentSimdOp.getReductionModAttr().getValue() ==
5468 ReductionModifier::inscan)
5469 return success();
5470 }
5471 return emitError("SCAN directive needs to be enclosed within a parent "
5472 "worksharing loop construct or SIMD construct with INSCAN "
5473 "reduction modifier");
5474}
5475
5476/// Verifies align clause in allocate directive
5477LogicalResult verifyAlignment(Operation &op,
5478 std::optional<uint64_t> alignment) {
5479 if (alignment.has_value()) {
5480 if ((alignment.value() != 0) && !llvm::has_single_bit(alignment.value()))
5481 return op.emitError()
5482 << "ALIGN value : " << alignment.value() << " must be power of 2";
5483 }
5484 return success();
5485}
5486
5487LogicalResult AllocateDirOp::verify() {
5488 return verifyAlignment(*getOperation(), getAlign());
5489}
5490
5491//===----------------------------------------------------------------------===//
5492// AllocSharedMemOp
5493//===----------------------------------------------------------------------===//
5494
5495LogicalResult AllocSharedMemOp::verify() {
5496 return verifyAlignment(*getOperation(), getMemAlignment());
5497}
5498
5499//===----------------------------------------------------------------------===//
5500// FreeSharedMemOp
5501//===----------------------------------------------------------------------===//
5502
5503LogicalResult FreeSharedMemOp::verify() {
5504 return verifyAlignment(*getOperation(), getMemAlignment());
5505}
5506
5507//===----------------------------------------------------------------------===//
5508// WorkdistributeOp
5509//===----------------------------------------------------------------------===//
5510
5511LogicalResult WorkdistributeOp::verify() {
5512 if (isCombined())
5513 return emitOpError() << "cannot be a non-innermost combined construct leaf";
5514
5515 // Check that region exists and is not empty
5516 Region &region = getRegion();
5517 if (region.empty())
5518 return emitOpError("region cannot be empty");
5519 // Verify single entry point.
5520 Block &entryBlock = region.front();
5521 if (entryBlock.empty())
5522 return emitOpError("region must contain a structured block");
5523 // Verify single exit point.
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");
5529 }
5530 hasTerminator = true;
5531 }
5532 }
5533 if (!hasTerminator) {
5534 return emitOpError("region must be terminated with omp.terminator");
5535 }
5536 auto walkResult = region.walk([&](Operation *op) -> WalkResult {
5537 // No implicit barrier at end
5538 if (isa<BarrierOp>(op)) {
5539 return emitOpError(
5540 "explicit barriers are not allowed in workdistribute region");
5541 }
5542 // Check for invalid nested constructs
5543 if (isa<ParallelOp>(op)) {
5544 return emitOpError(
5545 "nested parallel constructs not allowed in workdistribute");
5546 }
5547 if (isa<TeamsOp>(op)) {
5548 return emitOpError(
5549 "nested teams constructs not allowed in workdistribute");
5550 }
5551 return WalkResult::advance();
5552 });
5553 if (walkResult.wasInterrupted())
5554 return failure();
5555
5556 Operation *parentOp = (*this)->getParentOp();
5557 if (!llvm::dyn_cast<TeamsOp>(parentOp))
5558 return emitOpError("workdistribute must be nested under teams");
5559 return success();
5560}
5561
5562//===----------------------------------------------------------------------===//
5563// Declare simd [7.7]
5564//===----------------------------------------------------------------------===//
5565
5566LogicalResult DeclareSimdOp::verify() {
5567 // Must be nested inside a function-like op
5568 auto func =
5569 dyn_cast_if_present<mlir::FunctionOpInterface>((*this)->getParentOp());
5570 if (!func)
5571 return emitOpError() << "must be nested inside a function";
5572
5573 if (getInbranch() && getNotinbranch())
5574 return emitOpError("cannot have both 'inbranch' and 'notinbranch'");
5575
5576 if (failed(verifyLinearModifiers(*this, getLinearModifiers(), getLinearVars(),
5577 /*isDeclareSimd=*/true)))
5578 return failure();
5579
5580 return verifyAlignedClause(*this, getAlignments(), getAlignedVars());
5581}
5582
5583void DeclareSimdOp::build(OpBuilder &odsBuilder, OperationState &odsState,
5584 const DeclareSimdOperands &clauses) {
5585 MLIRContext *ctx = odsBuilder.getContext();
5586 DeclareSimdOp::build(odsBuilder, odsState, clauses.alignedVars,
5587 makeArrayAttr(ctx, clauses.alignments), clauses.inbranch,
5588 clauses.linearVars, clauses.linearStepVars,
5589 clauses.linearVarTypes, clauses.linearModifiers,
5590 clauses.notinbranch, clauses.simdlen,
5591 clauses.uniformVars);
5592}
5593
5594//===----------------------------------------------------------------------===//
5595// Parser and printer for Uniform Clause
5596//===----------------------------------------------------------------------===//
5597
5598/// uniform ::= `uniform` `(` uniform-list `)`
5599/// uniform-list := uniform-val (`,` uniform-val)*
5600/// uniform-val := ssa-id `:` type
5601static ParseResult
5604 SmallVectorImpl<Type> &uniformTypes) {
5605 return parser.parseCommaSeparatedList([&]() -> mlir::ParseResult {
5606 if (parser.parseOperand(uniformVars.emplace_back()) ||
5607 parser.parseColonType(uniformTypes.emplace_back()))
5608 return mlir::failure();
5609 return mlir::success();
5610 });
5611}
5612
5613/// Print Uniform Clauses
5615 ValueRange uniformVars, TypeRange uniformTypes) {
5616 for (unsigned i = 0; i < uniformVars.size(); ++i) {
5617 if (i != 0)
5618 p << ", ";
5619 p << uniformVars[i] << " : " << uniformTypes[i];
5620 }
5621}
5622
5623//===----------------------------------------------------------------------===//
5624// Parser and printer for Affinity Clause
5625//===----------------------------------------------------------------------===//
5626
5627static ParseResult parseAffinityClause(
5628 OpAsmParser &parser,
5631 SmallVectorImpl<Type> &iteratedTypes,
5632 SmallVectorImpl<Type> &affinityVarTypes) {
5633 if (failed(parseSplitIteratedList(
5634 parser, iterated, iteratedTypes, affinityVars, affinityVarTypes,
5635 /*parsePrefix=*/[&]() -> ParseResult { return success(); })))
5636 return failure();
5637 return success();
5638}
5639
5641 ValueRange iterated, ValueRange affinityVars,
5642 TypeRange iteratedTypes,
5643 TypeRange affinityVarTypes) {
5644 auto nop = [&](Value, Type) {};
5645 printSplitIteratedList(p, iterated, iteratedTypes, affinityVars,
5646 affinityVarTypes,
5647 /*plain prefix*/ nop,
5648 /*iterated prefix*/ nop);
5649}
5650
5651//===----------------------------------------------------------------------===//
5652// Parser, printer, and verifier for Iterator modifier
5653//===----------------------------------------------------------------------===//
5654
5655static ParseResult
5660 SmallVectorImpl<Type> &lbTypes,
5661 SmallVectorImpl<Type> &ubTypes,
5662 SmallVectorImpl<Type> &stepTypes) {
5663
5664 llvm::SMLoc ivLoc = parser.getCurrentLocation();
5666
5667 // Parse induction variables: %i : i32, %j : i32
5668 if (parser.parseCommaSeparatedList([&]() -> ParseResult {
5669 OpAsmParser::Argument &arg = ivArgs.emplace_back();
5670 if (parser.parseArgument(arg))
5671 return failure();
5672
5673 // Optional type, default to Index if not provided
5674 if (succeeded(parser.parseOptionalColon())) {
5675 if (parser.parseType(arg.type))
5676 return failure();
5677 } else {
5678 arg.type = parser.getBuilder().getIndexType();
5679 }
5680 return success();
5681 }))
5682 return failure();
5683
5684 // ) = (
5685 if (parser.parseRParen() || parser.parseEqual() || parser.parseLParen())
5686 return failure();
5687
5688 // Parse Ranges: (%lb to %ub step %st, ...)
5689 if (parser.parseCommaSeparatedList([&]() -> ParseResult {
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))
5694 return failure();
5695
5696 lbs.push_back(lb);
5697 ubs.push_back(ub);
5698 steps.push_back(st);
5699 return success();
5700 }))
5701 return failure();
5702
5703 if (parser.parseRParen())
5704 return failure();
5705
5706 if (ivArgs.size() != lbs.size())
5707 return parser.emitError(ivLoc)
5708 << "mismatch: " << ivArgs.size() << " variables but " << lbs.size()
5709 << " ranges";
5710
5711 for (auto &arg : ivArgs) {
5712 lbTypes.push_back(arg.type);
5713 ubTypes.push_back(arg.type);
5714 stepTypes.push_back(arg.type);
5715 }
5716
5717 return parser.parseRegion(region, ivArgs);
5718}
5719
5721 ValueRange lbs, ValueRange ubs,
5723 TypeRange) {
5724 Block &entry = region.front();
5725
5726 for (unsigned i = 0, e = entry.getNumArguments(); i < e; ++i) {
5727 if (i != 0)
5728 p << ", ";
5729 p.printRegionArgument(entry.getArgument(i));
5730 }
5731 p << ") = (";
5732
5733 // (%lb0 to %ub0 step %step0, %lb1 to %ub1 step %step1, ...)
5734 for (unsigned i = 0, e = lbs.size(); i < e; ++i) {
5735 if (i)
5736 p << ", ";
5737 p << lbs[i] << " to " << ubs[i] << " step " << steps[i];
5738 }
5739 p << ") ";
5740
5741 p.printRegion(region, /*printEntryBlockArgs=*/false,
5742 /*printBlockTerminators=*/true);
5743}
5744
5745LogicalResult IteratorOp::verify() {
5746 auto iteratedTy = llvm::dyn_cast<omp::IteratedType>(getIterated().getType());
5747 if (!iteratedTy)
5748 return emitOpError() << "result must be omp.iterated<entry_ty>";
5749
5750 for (auto [lb, ub, step] : llvm::zip_equal(
5751 getLoopLowerBounds(), getLoopUpperBounds(), getLoopSteps())) {
5752 if (matchPattern(step, m_Zero()))
5753 return emitOpError() << "loop step must not be zero";
5754
5755 IntegerAttr lbAttr;
5756 IntegerAttr ubAttr;
5757 IntegerAttr stepAttr;
5758 if (!matchPattern(lb, m_Constant(&lbAttr)) ||
5759 !matchPattern(ub, m_Constant(&ubAttr)) ||
5760 !matchPattern(step, m_Constant(&stepAttr)))
5761 continue;
5762
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";
5772 }
5773
5774 Block &b = getRegion().front();
5775 auto yield = llvm::dyn_cast<omp::YieldOp>(b.getTerminator());
5776
5777 if (!yield)
5778 return emitOpError() << "region must be terminated by omp.yield";
5779
5780 if (yield.getNumOperands() != 1)
5781 return emitOpError()
5782 << "omp.yield in omp.iterator region must yield exactly one value";
5783
5784 mlir::Type yieldedTy = yield.getOperand(0).getType();
5785 mlir::Type elemTy = iteratedTy.getElementType();
5786
5787 if (yieldedTy != elemTy)
5788 return emitOpError() << "omp.iterated element type (" << elemTy
5789 << ") does not match omp.yield operand type ("
5790 << yieldedTy << ")";
5791
5792 return success();
5793}
5794
5795//===----------------------------------------------------------------------===//
5796// GroupprivateOp
5797//===----------------------------------------------------------------------===//
5798
5799LogicalResult
5800GroupprivateOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
5801 auto *symbol = symbolTable.lookupNearestSymbolFrom(*this, getSymNameAttr());
5802 if (!symbol)
5803 return emitOpError() << "expected symbol reference '" << getSymName()
5804 << "' to point to a global variable";
5805
5806 if (isa<FunctionOpInterface>(symbol))
5807 return emitOpError() << "expected symbol reference '" << getSymName()
5808 << "' to point to a global variable, not a function";
5809
5810 return success();
5811}
5812
5813#define GET_ATTRDEF_CLASSES
5814#include "mlir/Dialect/OpenMP/OpenMPOpsAttributes.cpp.inc"
5815
5816#define GET_OP_CLASSES
5817#include "mlir/Dialect/OpenMP/OpenMPOps.cpp.inc"
5818
5819#define GET_TYPEDEF_CLASSES
5820#include "mlir/Dialect/OpenMP/OpenMPOpsTypes.cpp.inc"
return success()
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 &region, const Twine &name)
Definition EmitC.cpp:1523
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, ValueRange privateVars, TypeRange privateTypes, ArrayAttr privateSyms, UnitAttr privateNeedsBarrier)
static void printPrivateReductionRegion(OpAsmPrinter &p, Operation *op, Region &region, 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 &region, 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)
return success()
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 > &regionPrivateArgs, 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 &region, 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 &region, 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 &region, 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 &region, 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 &region, 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 > &copyprivateVars, SmallVectorImpl< Type > &copyprivateTypes, ArrayAttr &copyprivateSyms)
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,...
Definition SPIRVOps.cpp:229
static bool isUnique(It begin, It end)
Definition ShardOps.cpp:161
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.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
Definition Block.cpp:154
bool empty()
Definition Block.h:173
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
Operation & front()
Definition Block.h:178
SuccessorRange getSuccessors()
Definition Block.h:280
Operation & back()
Definition Block.h:177
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
bool mightHaveTerminator()
Return "true" if this block might have a terminator.
Definition Block.cpp:255
BlockArgListType getArguments()
Definition Block.h:112
iterator end()
Definition Block.h:169
iterator begin()
Definition Block.h:168
IntegerType getI64Type()
Definition Builders.cpp:73
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
MLIRContext * getContext() const
Definition Builders.h:56
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Definition Builders.h:101
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.
Definition Dominance.h:143
bool dominates(Operation *a, Operation *b) const
Return true if operation A dominates operation B, i.e.
Definition Dominance.h:161
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.
Definition MLIRContext.h:63
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
Definition Attributes.h:179
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region &region, 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.
Definition Builders.h:210
This class represents an operand of an operation.
Definition Value.h:254
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.
Definition ValueRange.h:44
type_range getType() const
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Definition Operation.h:237
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
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'.
Definition Operation.h:255
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
Region * getParentRegion()
Returns the region to which the instruction belongs.
Definition Operation.h:247
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
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.
Definition Region.h:26
Block & front()
Definition Region.h:65
BlockArgListType getArguments()
Definition Region.h:94
OpIterator op_begin()
Return iterators that walk the operations nested directly within this region.
Definition Region.h:178
bool isAncestor(Region *other)
Return true if this region is ancestor of the other region.
Definition Region.h:234
iterator_range< OpIterator > getOps()
Definition Region.h:180
bool empty()
Definition Region.h:60
unsigned getNumArguments()
Definition Region.h:136
Location getLoc()
Return a location for this region.
Definition Region.cpp:31
BlockArgument getArgument(unsigned i)
Definition Region.h:137
Operation * getParentOp()
Return the parent operation this region is attached to.
Definition Region.h:198
BlockListType & getBlocks()
Definition Region.h:45
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.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getType() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Definition Value.h:108
Type getType() const
Return the type of this value.
Definition Value.h:105
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
Definition Value.h:188
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult advance()
Definition WalkResult.h:47
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< bool > content)
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.
Definition UseDefLists.h:38
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
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.
Definition Utils.cpp:311
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
Definition LLVM.h:122
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.
Definition Matchers.h:442
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.
Definition Value.h:494
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
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...
Definition Utils.cpp:1380
detail::DenseArrayAttrImpl< bool > DenseBoolArrayAttr
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
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.