MLIR 24.0.0git
OpenACC.cpp
Go to the documentation of this file.
1//===- OpenACC.cpp - OpenACC MLIR Operations ------------------------------===//
2//
3// Part of the MLIR 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
16#include "mlir/IR/Builders.h"
18#include "mlir/IR/BuiltinOps.h"
21#include "mlir/IR/IRMapping.h"
22#include "mlir/IR/Matchers.h"
24#include "mlir/IR/SymbolTable.h"
25#include "mlir/Support/LLVM.h"
27#include "llvm/ADT/SmallSet.h"
28#include "llvm/ADT/TypeSwitch.h"
29#include "llvm/Support/LogicalResult.h"
30#include <variant>
31
32using namespace mlir;
33using namespace acc;
34
35#include "mlir/Dialect/OpenACC/OpenACCOpsDialect.cpp.inc"
36#include "mlir/Dialect/OpenACC/OpenACCOpsEnums.cpp.inc"
37#include "mlir/Dialect/OpenACC/OpenACCOpsInterfaces.cpp.inc"
38#include "mlir/Dialect/OpenACC/OpenACCTypeInterfaces.cpp.inc"
39#include "mlir/Dialect/OpenACCMPCommon/Interfaces/OpenACCMPOpsInterfaces.cpp.inc"
40
41namespace {
42
43static bool isScalarLikeType(Type type) {
44 return type.isIntOrIndexOrFloat() || isa<ComplexType>(type);
45}
46
47/// Helper function to attach the `VarName` attribute to an operation
48/// if a variable name is provided.
49static void attachVarNameAttr(Operation *op, OpBuilder &builder,
50 StringRef varName) {
51 if (!varName.empty()) {
52 auto varNameAttr = acc::VarNameAttr::get(builder.getContext(), varName);
54 }
55}
56
57template <typename T>
58struct MemRefPointerLikeModel
59 : public PointerLikeType::ExternalModel<MemRefPointerLikeModel<T>, T> {
60 Type getElementType(Type pointer) const {
61 return cast<T>(pointer).getElementType();
62 }
63
64 mlir::acc::VariableTypeCategory
65 getPointeeTypeCategory(Type pointer, TypedValue<PointerLikeType> varPtr,
66 Type varType) const {
67 if (auto mappableTy = dyn_cast<MappableType>(varType)) {
68 return mappableTy.getTypeCategory(varPtr);
69 }
70 auto memrefTy = cast<T>(pointer);
71 if (!memrefTy.hasRank()) {
72 // This memref is unranked - aka it could have any rank, including a
73 // rank of 0 which could mean scalar. For now, return uncategorized.
74 return mlir::acc::VariableTypeCategory::uncategorized;
75 }
76
77 if (memrefTy.getRank() == 0) {
78 if (isScalarLikeType(memrefTy.getElementType())) {
79 return mlir::acc::VariableTypeCategory::scalar;
80 }
81 // Zero-rank non-scalar - need further analysis to determine the type
82 // category. For now, return uncategorized.
83 return mlir::acc::VariableTypeCategory::uncategorized;
84 }
85
86 // It has a rank - must be an array.
87 assert(memrefTy.getRank() > 0 && "rank expected to be positive");
88 return mlir::acc::VariableTypeCategory::array;
89 }
90
91 mlir::Value genAllocate(Type pointer, OpBuilder &builder, Location loc,
92 StringRef varName, Type varType, Value originalVar,
93 bool &needsFree) const {
94 auto memrefTy = cast<MemRefType>(pointer);
95
96 // Check if this is a static memref (all dimensions are known) - if yes
97 // then we can generate an alloca operation.
98 if (memrefTy.hasStaticShape()) {
99 needsFree = false; // alloca doesn't need deallocation
100 auto allocaOp = memref::AllocaOp::create(builder, loc, memrefTy);
101 attachVarNameAttr(allocaOp, builder, varName);
102 return allocaOp.getResult();
103 }
104
105 // For dynamic memrefs, extract sizes from the original variable if
106 // provided. Otherwise they cannot be handled.
107 if (originalVar && originalVar.getType() == memrefTy &&
108 memrefTy.hasRank()) {
109 SmallVector<Value> dynamicSizes;
110 for (int64_t i = 0; i < memrefTy.getRank(); ++i) {
111 if (memrefTy.isDynamicDim(i)) {
112 // Extract the size of dimension i from the original variable
113 auto indexValue = arith::ConstantIndexOp::create(builder, loc, i);
114 auto dimSize =
115 memref::DimOp::create(builder, loc, originalVar, indexValue);
116 dynamicSizes.push_back(dimSize);
117 }
118 // Note: We only add dynamic sizes to the dynamicSizes array
119 // Static dimensions are handled automatically by AllocOp
120 }
121 needsFree = true; // alloc needs deallocation
122 auto allocOp =
123 memref::AllocOp::create(builder, loc, memrefTy, dynamicSizes);
124 attachVarNameAttr(allocOp, builder, varName);
125 return allocOp.getResult();
126 }
127
128 // TODO: Unranked not yet supported.
129 return {};
130 }
131
132 bool genFree(Type pointer, OpBuilder &builder, Location loc,
133 TypedValue<PointerLikeType> varToFree, Value allocRes,
134 Type varType) const {
135 if (auto memrefValue = dyn_cast<TypedValue<MemRefType>>(varToFree)) {
136 // Use allocRes if provided to determine the allocation type
137 Value valueToInspect = allocRes ? allocRes : memrefValue;
138
139 // Walk through casts to find the original allocation
140 Value currentValue = valueToInspect;
141 Operation *originalAlloc = nullptr;
142
143 // Follow the chain of operations to find the original allocation
144 // even if a casted result is provided.
145 while (currentValue) {
146 if (auto *definingOp = currentValue.getDefiningOp()) {
147 // Check if this is an allocation operation
148 if (isa<memref::AllocOp, memref::AllocaOp>(definingOp)) {
149 originalAlloc = definingOp;
150 break;
151 }
152
153 // Check if this is a cast operation we can look through
154 if (auto castOp = dyn_cast<memref::CastOp>(definingOp)) {
155 currentValue = castOp.getSource();
156 continue;
157 }
158
159 // Check for other cast-like operations
160 if (auto reinterpretCastOp =
161 dyn_cast<memref::ReinterpretCastOp>(definingOp)) {
162 currentValue = reinterpretCastOp.getSource();
163 continue;
164 }
165
166 // If we can't look through this operation, stop
167 break;
168 }
169 // This is a block argument or similar - can't trace further.
170 break;
171 }
172
173 if (originalAlloc) {
174 if (isa<memref::AllocaOp>(originalAlloc)) {
175 // This is an alloca - no dealloc needed, but return true (success)
176 return true;
177 }
178 if (isa<memref::AllocOp>(originalAlloc)) {
179 // This is an alloc - generate dealloc on varToFree
180 memref::DeallocOp::create(builder, loc, memrefValue);
181 return true;
182 }
183 }
184 }
185
186 return false;
187 }
188
189 bool genCopy(Type pointer, OpBuilder &builder, Location loc,
190 TypedValue<PointerLikeType> destination,
191 TypedValue<PointerLikeType> source, Type varType) const {
192 // Generate a copy operation between two memrefs
193 auto destMemref = dyn_cast_if_present<TypedValue<MemRefType>>(destination);
194 auto srcMemref = dyn_cast_if_present<TypedValue<MemRefType>>(source);
195
196 // As per memref documentation, source and destination must have same
197 // element type and shape in order to be compatible. We do not want to fail
198 // with an IR verification error - thus check that before generating the
199 // copy operation.
200 if (destMemref && srcMemref &&
201 destMemref.getType().getElementType() ==
202 srcMemref.getType().getElementType() &&
203 destMemref.getType().getShape() == srcMemref.getType().getShape()) {
204 memref::CopyOp::create(builder, loc, srcMemref, destMemref);
205 return true;
206 }
207
208 return false;
209 }
210
211 mlir::Value genLoad(Type pointer, OpBuilder &builder, Location loc,
213 Type valueType) const {
214 // Load from a memref - only valid for scalar memrefs (rank 0).
215 // This is because the address computation for memrefs is part of the load
216 // (and not computed separately), but the API does not have arguments for
217 // indexing.
218 auto memrefValue = dyn_cast_if_present<TypedValue<MemRefType>>(srcPtr);
219 if (!memrefValue)
220 return {};
221
222 auto memrefTy = memrefValue.getType();
223
224 // Only load from scalar memrefs (rank 0)
225 if (memrefTy.getRank() != 0)
226 return {};
227
228 return memref::LoadOp::create(builder, loc, memrefValue, ValueRange{});
229 }
230
231 bool genStore(Type pointer, OpBuilder &builder, Location loc,
232 Value valueToStore, TypedValue<PointerLikeType> destPtr) const {
233 // Store to a memref - only valid for scalar memrefs (rank 0)
234 // This is because the address computation for memrefs is part of the store
235 // (and not computed separately), but the API does not have arguments for
236 // indexing.
237 auto memrefValue = dyn_cast_if_present<TypedValue<MemRefType>>(destPtr);
238 if (!memrefValue)
239 return false;
240
241 auto memrefTy = memrefValue.getType();
242
243 // Only store to scalar memrefs (rank 0)
244 if (memrefTy.getRank() != 0)
245 return false;
246
247 memref::StoreOp::create(builder, loc, valueToStore, memrefValue);
248 return true;
249 }
250
251 Value genCast(Type, OpBuilder &builder, Location loc, Value value,
252 Type resultType) const {
253 if (value.getType() == resultType)
254 return value;
255
256 if (isa<BaseMemRefType>(value.getType()) &&
257 isa<BaseMemRefType>(resultType)) {
258 if (memref::CastOp::areCastCompatible(TypeRange(value.getType()),
259 TypeRange(resultType)))
260 return memref::CastOp::create(builder, loc, resultType, value);
261 if (memref::MemorySpaceCastOp::areCastCompatible(
262 TypeRange(value.getType()), TypeRange(resultType)))
263 return memref::MemorySpaceCastOp::create(builder, loc, resultType,
264 value);
265 }
266
267 // If one side is not a memref, try the other type's `PointerLikeType`
268 // implementation (since it may be an out-of-tree reference type that
269 // we cannot generate here).
270 if (auto resPtrLike = dyn_cast<PointerLikeType>(resultType))
271 if (!isa<BaseMemRefType>(resPtrLike))
272 if (Value v = resPtrLike.genCast(builder, loc, value, resultType))
273 return v;
274 if (auto valPtrLike = dyn_cast<PointerLikeType>(value.getType()))
275 if (!isa<BaseMemRefType>(valPtrLike))
276 if (Value v = valPtrLike.genCast(builder, loc, value, resultType))
277 return v;
278
279 return {};
280 }
281
282 bool isDeviceAccessible(Type pointer, Value var) const {
283 auto memrefTy = cast<T>(pointer);
284 Attribute memSpace = memrefTy.getMemorySpace();
285 return isa_and_nonnull<gpu::AddressSpaceAttr>(memSpace);
286 }
287
288 MemRefType getAsMemRefType(Type pointer, ModuleOp module) const {
289 (void)module;
290 return dyn_cast<MemRefType>(pointer);
291 }
292};
293
294struct LLVMPointerPointerLikeModel
295 : public PointerLikeType::ExternalModel<LLVMPointerPointerLikeModel,
296 LLVM::LLVMPointerType> {
297 Type getElementType(Type pointer) const { return Type(); }
298
299 mlir::Value genLoad(Type pointer, OpBuilder &builder, Location loc,
301 Type valueType) const {
302 // For LLVM pointers, we need the valueType to determine what to load
303 if (!valueType)
304 return {};
305
306 return LLVM::LoadOp::create(builder, loc, valueType, srcPtr);
307 }
308
309 bool genStore(Type pointer, OpBuilder &builder, Location loc,
310 Value valueToStore, TypedValue<PointerLikeType> destPtr) const {
311 LLVM::StoreOp::create(builder, loc, valueToStore, destPtr);
312 return true;
313 }
314
315 Value genCast(Type, OpBuilder &builder, Location loc, Value value,
316 Type resultType) const {
317 if (value.getType() == resultType)
318 return value;
319
320 auto srcPtrTy = dyn_cast<LLVM::LLVMPointerType>(value.getType());
321 auto dstPtrTy = dyn_cast<LLVM::LLVMPointerType>(resultType);
322 if (srcPtrTy && dstPtrTy) {
323 if (srcPtrTy.getAddressSpace() != dstPtrTy.getAddressSpace())
324 return LLVM::AddrSpaceCastOp::create(builder, loc, resultType, value);
325 return value;
326 }
327
328 if (srcPtrTy && isa<IntegerType>(resultType))
329 return LLVM::PtrToIntOp::create(builder, loc, resultType, value);
330
331 if (dstPtrTy) {
332 Value intVal = value;
333 if (isa<IndexType>(value.getType()))
334 intVal = arith::IndexCastUIOp::create(builder, loc,
335 builder.getI64Type(), value);
336 if (isa<IntegerType>(intVal.getType()))
337 return LLVM::IntToPtrOp::create(builder, loc, resultType, intVal);
338 }
339
340 if (auto resPtrLike = dyn_cast<PointerLikeType>(resultType))
341 if (!isa<LLVM::LLVMPointerType>(resPtrLike))
342 if (Value v = resPtrLike.genCast(builder, loc, value, resultType))
343 return v;
344 if (auto valPtrLike = dyn_cast<PointerLikeType>(value.getType()))
345 if (!isa<LLVM::LLVMPointerType>(valPtrLike))
346 if (Value v = valPtrLike.genCast(builder, loc, value, resultType))
347 return v;
348
349 return UnrealizedConversionCastOp::create(builder, loc,
350 TypeRange(resultType), value)
351 .getResult(0);
352 }
353};
354
355struct PrivateTypePointerLikeModel
356 : public PointerLikeType::ExternalModel<PrivateTypePointerLikeModel,
357 PrivateType> {
358 Type getElementType(Type type) const {
359 return cast<PrivateType>(type).getBaseTy();
360 }
361
362 Value genCast(Type, OpBuilder &builder, Location loc, Value value,
363 Type resultType) const {
364 if (value.getType() == resultType)
365 return value;
366 if (!isa<PointerLikeType>(resultType))
367 return {};
368 return UnwrapPrivateOp::create(builder, loc, resultType, value).getResult();
369 }
370
371 MemRefType getAsMemRefType(Type type, ModuleOp module) const {
372 Type baseTy = cast<PrivateType>(type).getBaseTy();
373 if (auto memrefTy = dyn_cast<MemRefType>(baseTy))
374 return memrefTy;
375 if (auto ptrLikeTy = dyn_cast<PointerLikeType>(baseTy))
376 return ptrLikeTy.getAsMemRefType(module);
377 return {};
378 }
379};
380
381struct MemrefAddressOfGlobalModel
382 : public AddressOfGlobalOpInterface::ExternalModel<
383 MemrefAddressOfGlobalModel, memref::GetGlobalOp> {
384 SymbolRefAttr getSymbol(Operation *op) const {
385 auto getGlobalOp = cast<memref::GetGlobalOp>(op);
386 return getGlobalOp.getNameAttr();
387 }
388};
389
390struct LLVMAddressOfGlobalModel
391 : public AddressOfGlobalOpInterface::ExternalModel<LLVMAddressOfGlobalModel,
392 LLVM::AddressOfOp> {
393 SymbolRefAttr getSymbol(Operation *op) const {
394 auto addressOfOp = cast<LLVM::AddressOfOp>(op);
395 return addressOfOp.getGlobalNameAttr();
396 }
397};
398
399struct MemrefGlobalVariableModel
400 : public GlobalVariableOpInterface::ExternalModel<MemrefGlobalVariableModel,
401 memref::GlobalOp> {
402 bool isConstant(Operation *op) const {
403 auto globalOp = cast<memref::GlobalOp>(op);
404 return globalOp.getConstant();
405 }
406
407 bool hasInitializer(Operation *op) const {
408 auto globalOp = cast<memref::GlobalOp>(op);
409 return globalOp.getInitialValue().has_value();
410 }
411
412 Region *getInitRegion(Operation *op) const {
413 // GlobalOp uses attributes for initialization, not regions
414 return nullptr;
415 }
416
417 bool isDeviceAccessible(Operation *op) const {
418 auto globalOp = cast<memref::GlobalOp>(op);
419 Attribute memSpace = globalOp.getType().getMemorySpace();
420 return isa_and_nonnull<gpu::AddressSpaceAttr>(memSpace);
421 }
422
423 bool isInDeviceMemory(Operation *op) const {
424 // A memref address space models storage that is physically resident on the
425 // device, so a device-accessible global is also in device memory. (There
426 // is no host-shared/migratable address space to exclude here.)
427 return isDeviceAccessible(op);
428 }
429
430 bool isCompilerGenerated(Operation *op) const { return false; }
431};
432
433struct GPULaunchOffloadRegionModel
434 : public acc::OffloadRegionOpInterface::ExternalModel<
435 GPULaunchOffloadRegionModel, gpu::LaunchOp> {
436 mlir::Region &getOffloadRegion(mlir::Operation *op) const {
437 return cast<gpu::LaunchOp>(op).getBody();
438 }
439};
440
441/// Helper function for any of the times we need to modify an ArrayAttr based on
442/// a device type list. Returns a new ArrayAttr with all of the
443/// existingDeviceTypes, plus the effective new ones(or an added none if hte new
444/// list is empty).
445mlir::ArrayAttr addDeviceTypeAffectedOperandHelper(
446 MLIRContext *context, mlir::ArrayAttr existingDeviceTypes,
447 llvm::ArrayRef<acc::DeviceType> newDeviceTypes) {
449 if (existingDeviceTypes)
450 llvm::copy(existingDeviceTypes, std::back_inserter(deviceTypes));
451
452 if (newDeviceTypes.empty())
453 deviceTypes.push_back(
454 acc::DeviceTypeAttr::get(context, acc::DeviceType::None));
455
456 for (DeviceType dt : newDeviceTypes)
457 deviceTypes.push_back(acc::DeviceTypeAttr::get(context, dt));
458
459 return mlir::ArrayAttr::get(context, deviceTypes);
460}
461
462/// Helper function for any of the times we need to add operands that are
463/// affected by a device type list. Returns a new ArrayAttr with all of the
464/// existingDeviceTypes, plus the effective new ones (or an added none, if the
465/// new list is empty). Additionally, adds the arguments to the argCollection
466/// the correct number of times. This will also update a 'segments' array, even
467/// if it won't be used.
468mlir::ArrayAttr addDeviceTypeAffectedOperandHelper(
469 MLIRContext *context, mlir::ArrayAttr existingDeviceTypes,
470 llvm::ArrayRef<acc::DeviceType> newDeviceTypes, mlir::ValueRange arguments,
471 mlir::MutableOperandRange argCollection,
472 llvm::SmallVector<int32_t> &segments) {
474 if (existingDeviceTypes)
475 llvm::copy(existingDeviceTypes, std::back_inserter(deviceTypes));
476
477 if (newDeviceTypes.empty()) {
478 argCollection.append(arguments);
479 segments.push_back(arguments.size());
480 deviceTypes.push_back(
481 acc::DeviceTypeAttr::get(context, acc::DeviceType::None));
482 }
483
484 for (DeviceType dt : newDeviceTypes) {
485 argCollection.append(arguments);
486 segments.push_back(arguments.size());
487 deviceTypes.push_back(acc::DeviceTypeAttr::get(context, dt));
488 }
489
490 return mlir::ArrayAttr::get(context, deviceTypes);
491}
492
493/// Overload for when the 'segments' aren't needed.
494mlir::ArrayAttr addDeviceTypeAffectedOperandHelper(
495 MLIRContext *context, mlir::ArrayAttr existingDeviceTypes,
496 llvm::ArrayRef<acc::DeviceType> newDeviceTypes, mlir::ValueRange arguments,
497 mlir::MutableOperandRange argCollection) {
499 return addDeviceTypeAffectedOperandHelper(context, existingDeviceTypes,
500 newDeviceTypes, arguments,
501 argCollection, segments);
502}
503} // namespace
504
505//===----------------------------------------------------------------------===//
506// OpenACC operations
507//===----------------------------------------------------------------------===//
508
509void OpenACCDialect::initialize() {
510 addOperations<
511#define GET_OP_LIST
512#include "mlir/Dialect/OpenACC/OpenACCOps.cpp.inc"
513 >();
514 addAttributes<
515#define GET_ATTRDEF_LIST
516#include "mlir/Dialect/OpenACC/OpenACCOpsAttributes.cpp.inc"
517 >();
518 addTypes<
519#define GET_TYPEDEF_LIST
520#include "mlir/Dialect/OpenACC/OpenACCOpsTypes.cpp.inc"
521 >();
522
523 // By attaching interfaces here, we make the OpenACC dialect dependent on
524 // the other dialects. This is probably better than having dialects like LLVM
525 // and memref be dependent on OpenACC.
526 MemRefType::attachInterface<MemRefPointerLikeModel<MemRefType>>(
527 *getContext());
528 UnrankedMemRefType::attachInterface<
529 MemRefPointerLikeModel<UnrankedMemRefType>>(*getContext());
530 LLVM::LLVMPointerType::attachInterface<LLVMPointerPointerLikeModel>(
531 *getContext());
532 PrivateType::attachInterface<PrivateTypePointerLikeModel>(*getContext());
533
534 // Attach operation interfaces
535 memref::GetGlobalOp::attachInterface<MemrefAddressOfGlobalModel>(
536 *getContext());
537 LLVM::AddressOfOp::attachInterface<LLVMAddressOfGlobalModel>(*getContext());
538 memref::GlobalOp::attachInterface<MemrefGlobalVariableModel>(*getContext());
539 gpu::LaunchOp::attachInterface<GPULaunchOffloadRegionModel>(*getContext());
540}
541
542//===----------------------------------------------------------------------===//
543// RegionBranchOpInterface for acc.kernels / acc.parallel / acc.serial /
544// acc.kernel_environment / acc.data / acc.host_data / acc.loop
545//===----------------------------------------------------------------------===//
546
547/// Generic helper for single-region OpenACC ops that execute their body once
548/// and then continue after the operation with their results (if any).
549static void
551 RegionBranchPoint point,
553 if (point.isParent()) {
554 regions.push_back(RegionSuccessor(&region));
555 return;
556 }
557
558 regions.push_back(RegionSuccessor(op));
559}
560
562 RegionSuccessor successor) {
563 return successor.isOperation() ? ValueRange(op->getResults()) : ValueRange();
564}
565
566void KernelsOp::getSuccessorRegions(RegionBranchPoint point,
568 getSingleRegionOpSuccessorRegions(getOperation(), getRegion(), point,
569 regions);
570}
571
572ValueRange KernelsOp::getSuccessorInputs(RegionSuccessor successor) {
573 return getSingleRegionSuccessorInputs(getOperation(), successor);
574}
575
576void ParallelOp::getSuccessorRegions(
578 getSingleRegionOpSuccessorRegions(getOperation(), getRegion(), point,
579 regions);
580}
581
582ValueRange ParallelOp::getSuccessorInputs(RegionSuccessor successor) {
583 return getSingleRegionSuccessorInputs(getOperation(), successor);
584}
585
586void SerialOp::getSuccessorRegions(RegionBranchPoint point,
588 getSingleRegionOpSuccessorRegions(getOperation(), getRegion(), point,
589 regions);
590}
591
592ValueRange SerialOp::getSuccessorInputs(RegionSuccessor successor) {
593 return getSingleRegionSuccessorInputs(getOperation(), successor);
594}
595
596void DataOp::getSuccessorRegions(RegionBranchPoint point,
598 getSingleRegionOpSuccessorRegions(getOperation(), getRegion(), point,
599 regions);
600}
601
602ValueRange DataOp::getSuccessorInputs(RegionSuccessor successor) {
603 return getSingleRegionSuccessorInputs(getOperation(), successor);
604}
605
606void HostDataOp::getSuccessorRegions(
608 getSingleRegionOpSuccessorRegions(getOperation(), getRegion(), point,
609 regions);
610}
611
612ValueRange HostDataOp::getSuccessorInputs(RegionSuccessor successor) {
613 return getSingleRegionSuccessorInputs(getOperation(), successor);
614}
615
616/// Whether the body of a structured `acc.loop` is proven to run. This decides
617/// which edges out of the parent are feasible; the edges out of the region are
618/// unaffected.
619enum class BodyExecution {
620 /// The body runs at least once, so the parent cannot bypass the region.
622 /// The body never runs, so the parent cannot enter the region.
624 /// Neither could be proven, so the parent may do either.
626};
627
628/// Prove whether the body of `loopOp` runs. A counted dimension whose entry
629/// test already fails at its lower bound runs exactly zero times, so a single
630/// comparison decides both `Always` and `Never`. Bounds that are not constant
631/// prove nothing.
632static BodyExecution getBodyExecution(LoopOp loopOp) {
633 // A container-like loop describes its iteration space inside the region.
634 if (loopOp.isContainerLike())
636
637 // The verifier guarantees one lower bound, upper bound and step per
638 // dimension.
639 ValueRange lbs = loopOp.getLowerbound();
640 ValueRange ubs = loopOp.getUpperbound();
641 ValueRange steps = loopOp.getStep();
642
644 for (unsigned i = 0, e = lbs.size(); i < e; ++i) {
645 std::optional<int64_t> lb = getConstantIntValue(lbs[i]);
646 std::optional<int64_t> ub = getConstantIntValue(ubs[i]);
647 std::optional<int64_t> step = getConstantIntValue(steps[i]);
648 // An unknown bound cannot be tested, and a zero step either spins forever
649 // or never starts. Neither proves `Always`, but a later dimension may
650 // still prove the nest empty: `(0 to %n)` collapsed with `(0 to 0)`.
651 if (!lb || !ub || !step || *step == 0) {
653 continue;
654 }
655
656 // The entry test at the lower bound. A descending dimension compares
657 // against its bound the other way round. The attribute is absent when
658 // every dimension is exclusive as in `scf.for`, and the verifier otherwise
659 // guarantees one entry per dimension.
660 std::optional<ArrayRef<bool>> inclusiveUbs =
661 loopOp.getInclusiveUpperbound();
662 bool inclusiveUb = inclusiveUbs && (*inclusiveUbs)[i];
663 assert(*step != 0 && "zero step should have been filtered out");
664 bool runsOnce = *step > 0 ? (inclusiveUb ? *lb <= *ub : *lb < *ub)
665 : (inclusiveUb ? *lb >= *ub : *lb > *ub);
666 // The dimensions are iterated as a nest, so one empty dimension empties
667 // the whole nest whatever the others do, while the body runs only if every
668 // dimension runs.
669 if (!runsOnce)
671 }
672
673 // No dimension was empty, so the body runs unless some dimension was
674 // unknown.
675 return result;
676}
677
678void LoopOp::getSuccessorRegions(RegionBranchPoint point,
680 // Unstructured loops: the body may contain arbitrary CFG and early exits.
681 // At the RegionBranch level, only model entry into the body and exit to the
682 // parent; any backedges are represented inside the region CFG.
683 if (getUnstructured()) {
684 if (point.isParent()) {
685 regions.push_back(RegionSuccessor(&getRegion()));
686 return;
687 }
688 regions.push_back(RegionSuccessor(getOperation()));
689 return;
690 }
691
692 // Structured loops: model a loop-shaped region graph similar to scf.for,
693 // minus the entry edge the loop is proven not to take.
694 if (point.isParent()) {
695 switch (getBodyExecution(*this)) {
697 regions.push_back(RegionSuccessor(&getRegion()));
698 return;
700 regions.push_back(RegionSuccessor(getOperation()));
701 return;
703 break;
704 }
705 }
706
707 regions.push_back(RegionSuccessor(&getRegion()));
708 regions.push_back(RegionSuccessor(getOperation()));
709}
710
711ValueRange LoopOp::getSuccessorInputs(RegionSuccessor successor) {
712 return getSingleRegionSuccessorInputs(getOperation(), successor);
713}
714
715//===----------------------------------------------------------------------===//
716// RegionBranchTerminatorOpInterface
717//===----------------------------------------------------------------------===//
718
720TerminatorOp::getMutableSuccessorOperands(RegionSuccessor /*point*/) {
721 // `acc.terminator` does not forward operands.
722 return MutableOperandRange(getOperation(), /*start=*/0, /*length=*/0);
723}
724
725//===----------------------------------------------------------------------===//
726// device_type support helpers
727//===----------------------------------------------------------------------===//
728
729static bool hasDeviceTypeValues(std::optional<mlir::ArrayAttr> arrayAttr) {
730 return arrayAttr && *arrayAttr && arrayAttr->size() > 0;
731}
732
733static bool hasDeviceType(std::optional<mlir::ArrayAttr> arrayAttr,
734 mlir::acc::DeviceType deviceType) {
735 if (!hasDeviceTypeValues(arrayAttr))
736 return false;
737
738 for (auto attr : *arrayAttr) {
739 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
740 if (deviceTypeAttr.getValue() == deviceType)
741 return true;
742 }
743
744 return false;
745}
746
748 std::optional<mlir::ArrayAttr> deviceTypes) {
749 if (!hasDeviceTypeValues(deviceTypes))
750 return;
751
752 p << "[";
753 llvm::interleaveComma(*deviceTypes, p,
754 [&](mlir::Attribute attr) { p << attr; });
755 p << "]";
756}
757
758static std::optional<unsigned> findSegment(ArrayAttr segments,
759 mlir::acc::DeviceType deviceType) {
760 unsigned segmentIdx = 0;
761 for (auto attr : segments) {
762 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
763 if (deviceTypeAttr.getValue() == deviceType)
764 return std::make_optional(segmentIdx);
765 ++segmentIdx;
766 }
767 return std::nullopt;
768}
769
771getValuesFromSegments(std::optional<mlir::ArrayAttr> arrayAttr,
773 std::optional<llvm::ArrayRef<int32_t>> segments,
774 mlir::acc::DeviceType deviceType) {
775 if (!arrayAttr)
776 return range.take_front(0);
777 if (auto pos = findSegment(*arrayAttr, deviceType)) {
778 int32_t nbOperandsBefore = 0;
779 for (unsigned i = 0; i < *pos; ++i)
780 nbOperandsBefore += (*segments)[i];
781 return range.drop_front(nbOperandsBefore).take_front((*segments)[*pos]);
782 }
783 return range.take_front(0);
784}
785
786static mlir::Value
787getWaitDevnumValue(std::optional<mlir::ArrayAttr> deviceTypeAttr,
789 std::optional<llvm::ArrayRef<int32_t>> segments,
790 std::optional<mlir::ArrayAttr> hasWaitDevnum,
791 mlir::acc::DeviceType deviceType) {
792 if (!hasDeviceTypeValues(deviceTypeAttr))
793 return {};
794 if (auto pos = findSegment(*deviceTypeAttr, deviceType)) {
796 auto boolAttr = mlir::dyn_cast<mlir::BoolAttr>((*hasWaitDevnum)[*pos]);
797 if (boolAttr && boolAttr.getValue())
798 return getValuesFromSegments(deviceTypeAttr, operands, segments,
799 deviceType)
800 .front();
801 }
802 }
803 return {};
804}
805
807getWaitValuesWithoutDevnum(std::optional<mlir::ArrayAttr> deviceTypeAttr,
809 std::optional<llvm::ArrayRef<int32_t>> segments,
810 std::optional<mlir::ArrayAttr> hasWaitDevnum,
811 mlir::acc::DeviceType deviceType) {
812 auto range =
813 getValuesFromSegments(deviceTypeAttr, operands, segments, deviceType);
814 if (range.empty())
815 return range;
816 if (auto pos = findSegment(*deviceTypeAttr, deviceType)) {
818 auto boolAttr = mlir::dyn_cast<mlir::BoolAttr>((*hasWaitDevnum)[*pos]);
819 if (boolAttr.getValue())
820 return range.drop_front(1); // first value is devnum
821 }
822 }
823 return range;
824}
825
826template <typename Op>
827static LogicalResult checkWaitAndAsyncConflict(Op op) {
828 for (uint32_t dtypeInt = 0; dtypeInt != acc::getMaxEnumValForDeviceType();
829 ++dtypeInt) {
830 auto dtype = static_cast<acc::DeviceType>(dtypeInt);
831
832 // The asyncOnly attribute represent the async clause without value.
833 // Therefore the attribute and operand cannot appear at the same time.
834 if (hasDeviceType(op.getAsyncOperandsDeviceType(), dtype) &&
835 op.hasAsyncOnly(dtype))
836 return op.emitError(
837 "asyncOnly attribute cannot appear with asyncOperand");
838
839 // The wait attribute represent the wait clause without values. Therefore
840 // the attribute and operands cannot appear at the same time.
841 if (hasDeviceType(op.getWaitOperandsDeviceType(), dtype) &&
842 op.hasWaitOnly(dtype))
843 return op.emitError("wait attribute cannot appear with waitOperands");
844 }
845 return success();
846}
847
848template <typename Op>
849static LogicalResult checkVarAndVarType(Op op) {
850 if (!op.getVar())
851 return op.emitError("must have var operand");
852
853 // A variable must have a type that is either pointer-like or mappable.
854 if (!mlir::isa<mlir::acc::PointerLikeType>(op.getVar().getType()) &&
855 !mlir::isa<mlir::acc::MappableType>(op.getVar().getType()))
856 return op.emitError("var must be mappable or pointer-like");
857
858 // When it is a pointer-like type, the varType must capture the target type.
859 if (mlir::isa<mlir::acc::PointerLikeType>(op.getVar().getType()) &&
860 op.getVarType() == op.getVar().getType())
861 return op.emitError("varType must capture the element type of var");
862
863 return success();
864}
865
866template <typename Op>
867static LogicalResult checkVarAndAccVar(Op op) {
868 if (op.getVar().getType() != op.getAccVar().getType())
869 return op.emitError("input and output types must match");
870
871 return success();
872}
873
874template <typename Op>
875static LogicalResult checkNoModifier(Op op) {
876 if (op.getModifiers() != acc::DataClauseModifier::none)
877 return op.emitError("no data clause modifiers are allowed");
878 return success();
879}
880
881template <typename Op>
882static LogicalResult
883checkValidModifier(Op op, acc::DataClauseModifier validModifiers) {
884 if (acc::bitEnumContainsAny(op.getModifiers(), ~validModifiers))
885 return op.emitError(
886 "invalid data clause modifiers: " +
887 acc::stringifyDataClauseModifier(op.getModifiers() & ~validModifiers));
888
889 return success();
890}
891
892template <typename OpT, typename RecipeOpT>
893static LogicalResult checkRecipe(OpT op, llvm::StringRef operandName) {
894 // Mappable types do not need a recipe because it is possible to generate one
895 // from its API. Reject reductions though because no API is available for them
896 // at this time.
897 if (mlir::acc::isMappableType(op.getVar().getType()) &&
898 !std::is_same_v<OpT, acc::ReductionOp>)
899 return success();
900
901 mlir::SymbolRefAttr operandRecipe = op.getRecipeAttr();
902 if (!operandRecipe)
903 return op->emitOpError() << "recipe expected for " << operandName;
904
905 auto decl =
907 if (!decl)
908 return op->emitOpError()
909 << "expected symbol reference " << operandRecipe << " to point to a "
910 << operandName << " declaration";
911 return success();
912}
913
914static ParseResult parseVar(mlir::OpAsmParser &parser,
916 // Either `var` or `varPtr` keyword is required.
917 if (failed(parser.parseOptionalKeyword("varPtr"))) {
918 if (failed(parser.parseKeyword("var")))
919 return failure();
920 }
921 if (failed(parser.parseLParen()))
922 return failure();
923 if (failed(parser.parseOperand(var)))
924 return failure();
925
926 return success();
927}
928
930 mlir::Value var) {
931 if (mlir::isa<mlir::acc::PointerLikeType>(var.getType()))
932 p << "varPtr(";
933 else
934 p << "var(";
935 p.printOperand(var);
936}
937
938static ParseResult parseAccVar(mlir::OpAsmParser &parser,
940 mlir::Type &accVarType) {
941 // Either `accVar` or `accPtr` keyword is required.
942 if (failed(parser.parseOptionalKeyword("accPtr"))) {
943 if (failed(parser.parseKeyword("accVar")))
944 return failure();
945 }
946 if (failed(parser.parseLParen()))
947 return failure();
948 if (failed(parser.parseOperand(var)))
949 return failure();
950 if (failed(parser.parseColon()))
951 return failure();
952 if (failed(parser.parseType(accVarType)))
953 return failure();
954 if (failed(parser.parseRParen()))
955 return failure();
956
957 return success();
958}
959
961 mlir::Value accVar, mlir::Type accVarType) {
962 if (mlir::isa<mlir::acc::PointerLikeType>(accVar.getType()))
963 p << "accPtr(";
964 else
965 p << "accVar(";
966 p.printOperand(accVar);
967 p << " : ";
968 p.printType(accVarType);
969 p << ")";
970}
971
972static ParseResult parseVarPtrType(mlir::OpAsmParser &parser,
973 mlir::Type &varPtrType,
974 mlir::TypeAttr &varTypeAttr) {
975 if (failed(parser.parseType(varPtrType)))
976 return failure();
977 if (failed(parser.parseRParen()))
978 return failure();
979
980 if (succeeded(parser.parseOptionalKeyword("varType"))) {
981 if (failed(parser.parseLParen()))
982 return failure();
983 mlir::Type varType;
984 if (failed(parser.parseType(varType)))
985 return failure();
986 varTypeAttr = mlir::TypeAttr::get(varType);
987 if (failed(parser.parseRParen()))
988 return failure();
989 } else {
990 // Set `varType` from the element type of the type of `varPtr`.
991 if (auto ptrTy = dyn_cast<acc::PointerLikeType>(varPtrType)) {
992 Type elementType = ptrTy.getElementType();
993 // Opaque pointers (e.g. !llvm.ptr) have no element type; fall back to
994 // using varPtrType itself so that the attribute is always valid.
995 varTypeAttr = mlir::TypeAttr::get(elementType ? elementType : varPtrType);
996 } else {
997 varTypeAttr = mlir::TypeAttr::get(varPtrType);
998 }
999 }
1000
1001 return success();
1002}
1003
1005 mlir::Type varPtrType, mlir::TypeAttr varTypeAttr) {
1006 p.printType(varPtrType);
1007 p << ")";
1008
1009 // Print the `varType` only if it differs from the element type of
1010 // `varPtr`'s type.
1011 mlir::Type varType = varTypeAttr.getValue();
1012 mlir::Type typeToCheckAgainst =
1013 mlir::isa<mlir::acc::PointerLikeType>(varPtrType)
1014 ? mlir::cast<mlir::acc::PointerLikeType>(varPtrType).getElementType()
1015 : varPtrType;
1016 // Opaque pointers (e.g. !llvm.ptr) have no element type; use varPtrType as
1017 // the baseline so that the inferred varType is not redundantly printed.
1018 if (!typeToCheckAgainst)
1019 typeToCheckAgainst = varPtrType;
1020 if (typeToCheckAgainst != varType) {
1021 p << " varType(";
1022 p.printType(varType);
1023 p << ")";
1024 }
1025}
1026
1027// A location cannot be parsed through the generic attribute directive: the
1028// generated parser needs a concrete attribute class, while any of the location
1029// attributes may appear here.
1030static ParseResult parseSourceLocation(mlir::OpAsmParser &parser,
1031 mlir::LocationAttr &locAttr) {
1032 llvm::SMLoc attrLoc = parser.getCurrentLocation();
1033 mlir::Attribute attr;
1034 if (failed(parser.parseAttribute(attr)))
1035 return failure();
1036 locAttr = mlir::dyn_cast<mlir::LocationAttr>(attr);
1037 if (!locAttr)
1038 return parser.emitError(attrLoc, "expected location attribute");
1039 return success();
1040}
1041
1043 mlir::LocationAttr locAttr) {
1044 p.printAttribute(locAttr);
1045}
1046
1047static ParseResult parseRecipeSym(mlir::OpAsmParser &parser,
1048 mlir::SymbolRefAttr &recipeAttr) {
1049 if (failed(parser.parseAttribute(recipeAttr)))
1050 return failure();
1051 return success();
1052}
1053
1055 mlir::SymbolRefAttr recipeAttr) {
1056 p << recipeAttr;
1057}
1058
1061 return parser.parseAttribute(attr);
1062}
1063
1066 p << attr;
1067}
1068
1069static ParseResult parseArrayAttr(mlir::OpAsmParser &parser,
1070 mlir::ArrayAttr &attr) {
1071 return parser.parseAttribute(attr);
1072}
1073
1075 mlir::ArrayAttr attr) {
1076 p << attr;
1077}
1078
1079//===----------------------------------------------------------------------===//
1080// DataBoundsOp
1081//===----------------------------------------------------------------------===//
1082LogicalResult acc::DataBoundsOp::verify() {
1083 auto extent = getExtent();
1084 auto upperbound = getUpperbound();
1085 if (!extent && !upperbound)
1086 return emitError("expected extent or upperbound.");
1087 return success();
1088}
1089
1090//===----------------------------------------------------------------------===//
1091// PrivateOp
1092//===----------------------------------------------------------------------===//
1093LogicalResult acc::PrivateOp::verify() {
1094 if (getDataClause() != acc::DataClause::acc_private)
1095 return emitError(
1096 "data clause associated with private operation must match its intent");
1097 if (failed(checkVarAndVarType(*this)))
1098 return failure();
1099 if (failed(checkNoModifier(*this)))
1100 return failure();
1101 if (failed(
1103 return failure();
1104 return success();
1105}
1106
1107//===----------------------------------------------------------------------===//
1108// FirstprivateOp
1109//===----------------------------------------------------------------------===//
1110LogicalResult acc::FirstprivateOp::verify() {
1111 if (getDataClause() != acc::DataClause::acc_firstprivate)
1112 return emitError("data clause associated with firstprivate operation must "
1113 "match its intent");
1114 if (failed(checkVarAndVarType(*this)))
1115 return failure();
1116 if (failed(checkNoModifier(*this)))
1117 return failure();
1119 *this, "firstprivate")))
1120 return failure();
1121 return success();
1122}
1123
1124//===----------------------------------------------------------------------===//
1125// ReductionOp
1126//===----------------------------------------------------------------------===//
1127LogicalResult acc::ReductionOp::verify() {
1128 if (getDataClause() != acc::DataClause::acc_reduction)
1129 return emitError("data clause associated with reduction operation must "
1130 "match its intent");
1131 if (failed(checkVarAndVarType(*this)))
1132 return failure();
1133 if (failed(checkNoModifier(*this)))
1134 return failure();
1136 *this, "reduction")))
1137 return failure();
1138 return success();
1139}
1140
1141//===----------------------------------------------------------------------===//
1142// DevicePtrOp
1143//===----------------------------------------------------------------------===//
1144LogicalResult acc::DevicePtrOp::verify() {
1145 if (getDataClause() != acc::DataClause::acc_deviceptr)
1146 return emitError("data clause associated with deviceptr operation must "
1147 "match its intent");
1148 if (failed(checkVarAndVarType(*this)))
1149 return failure();
1150 if (failed(checkVarAndAccVar(*this)))
1151 return failure();
1152 if (failed(checkNoModifier(*this)))
1153 return failure();
1154 return success();
1155}
1156
1157//===----------------------------------------------------------------------===//
1158// PresentOp
1159//===----------------------------------------------------------------------===//
1160LogicalResult acc::PresentOp::verify() {
1161 if (getDataClause() != acc::DataClause::acc_present)
1162 return emitError(
1163 "data clause associated with present operation must match its intent");
1164 if (failed(checkVarAndVarType(*this)))
1165 return failure();
1166 if (failed(checkVarAndAccVar(*this)))
1167 return failure();
1168 if (failed(checkNoModifier(*this)))
1169 return failure();
1170 return success();
1171}
1172
1173//===----------------------------------------------------------------------===//
1174// CopyinOp
1175//===----------------------------------------------------------------------===//
1176LogicalResult acc::CopyinOp::verify() {
1177 // Test for all clauses this operation can be decomposed from:
1178 if (!getImplicit() && getDataClause() != acc::DataClause::acc_copyin &&
1179 getDataClause() != acc::DataClause::acc_copyin_readonly &&
1180 getDataClause() != acc::DataClause::acc_copy &&
1181 getDataClause() != acc::DataClause::acc_reduction)
1182 return emitError(
1183 "data clause associated with copyin operation must match its intent"
1184 " or specify original clause this operation was decomposed from");
1185 if (failed(checkVarAndVarType(*this)))
1186 return failure();
1187 if (failed(checkVarAndAccVar(*this)))
1188 return failure();
1189 if (failed(checkValidModifier(*this, acc::DataClauseModifier::readonly |
1190 acc::DataClauseModifier::always |
1191 acc::DataClauseModifier::capture)))
1192 return failure();
1193 return success();
1194}
1195
1196bool acc::CopyinOp::isCopyinReadonly() {
1197 return getDataClause() == acc::DataClause::acc_copyin_readonly ||
1198 acc::bitEnumContainsAny(getModifiers(),
1199 acc::DataClauseModifier::readonly);
1200}
1201
1202//===----------------------------------------------------------------------===//
1203// CreateOp
1204//===----------------------------------------------------------------------===//
1205LogicalResult acc::CreateOp::verify() {
1206 // Test for all clauses this operation can be decomposed from:
1207 if (getDataClause() != acc::DataClause::acc_create &&
1208 getDataClause() != acc::DataClause::acc_create_zero &&
1209 getDataClause() != acc::DataClause::acc_copyout &&
1210 getDataClause() != acc::DataClause::acc_copyout_zero)
1211 return emitError(
1212 "data clause associated with create operation must match its intent"
1213 " or specify original clause this operation was decomposed from");
1214 if (failed(checkVarAndVarType(*this)))
1215 return failure();
1216 if (failed(checkVarAndAccVar(*this)))
1217 return failure();
1218 // this op is the entry part of copyout, so it also needs to allow all
1219 // modifiers allowed on copyout.
1220 if (failed(checkValidModifier(*this, acc::DataClauseModifier::zero |
1221 acc::DataClauseModifier::always |
1222 acc::DataClauseModifier::capture)))
1223 return failure();
1224 return success();
1225}
1226
1227bool acc::CreateOp::isCreateZero() {
1228 // The zero modifier is encoded in the data clause.
1229 return getDataClause() == acc::DataClause::acc_create_zero ||
1230 getDataClause() == acc::DataClause::acc_copyout_zero ||
1231 acc::bitEnumContainsAny(getModifiers(), acc::DataClauseModifier::zero);
1232}
1233
1234//===----------------------------------------------------------------------===//
1235// NoCreateOp
1236//===----------------------------------------------------------------------===//
1237LogicalResult acc::NoCreateOp::verify() {
1238 if (getDataClause() != acc::DataClause::acc_no_create)
1239 return emitError("data clause associated with no_create operation must "
1240 "match its intent");
1241 if (failed(checkVarAndVarType(*this)))
1242 return failure();
1243 if (failed(checkVarAndAccVar(*this)))
1244 return failure();
1245 if (failed(checkNoModifier(*this)))
1246 return failure();
1247 return success();
1248}
1249
1250//===----------------------------------------------------------------------===//
1251// AttachOp
1252//===----------------------------------------------------------------------===//
1253LogicalResult acc::AttachOp::verify() {
1254 if (getDataClause() != acc::DataClause::acc_attach)
1255 return emitError(
1256 "data clause associated with attach operation must match its intent");
1257 if (failed(checkVarAndVarType(*this)))
1258 return failure();
1259 if (failed(checkVarAndAccVar(*this)))
1260 return failure();
1261 if (failed(checkNoModifier(*this)))
1262 return failure();
1263 return success();
1264}
1265
1266//===----------------------------------------------------------------------===//
1267// DeclareDeviceResidentOp
1268//===----------------------------------------------------------------------===//
1269
1270LogicalResult acc::DeclareDeviceResidentOp::verify() {
1271 if (getDataClause() != acc::DataClause::acc_declare_device_resident)
1272 return emitError("data clause associated with device_resident operation "
1273 "must match its intent");
1274 if (failed(checkVarAndVarType(*this)))
1275 return failure();
1276 if (failed(checkVarAndAccVar(*this)))
1277 return failure();
1278 if (failed(checkNoModifier(*this)))
1279 return failure();
1280 return success();
1281}
1282
1283//===----------------------------------------------------------------------===//
1284// DeclareLinkOp
1285//===----------------------------------------------------------------------===//
1286
1287LogicalResult acc::DeclareLinkOp::verify() {
1288 if (getDataClause() != acc::DataClause::acc_declare_link)
1289 return emitError(
1290 "data clause associated with link operation must match its intent");
1291 if (failed(checkVarAndVarType(*this)))
1292 return failure();
1293 if (failed(checkVarAndAccVar(*this)))
1294 return failure();
1295 if (failed(checkNoModifier(*this)))
1296 return failure();
1297 return success();
1298}
1299
1300//===----------------------------------------------------------------------===//
1301// CopyoutOp
1302//===----------------------------------------------------------------------===//
1303LogicalResult acc::CopyoutOp::verify() {
1304 // Test for all clauses this operation can be decomposed from:
1305 if (getDataClause() != acc::DataClause::acc_copyout &&
1306 getDataClause() != acc::DataClause::acc_copyout_zero &&
1307 getDataClause() != acc::DataClause::acc_copy &&
1308 getDataClause() != acc::DataClause::acc_reduction)
1309 return emitError(
1310 "data clause associated with copyout operation must match its intent"
1311 " or specify original clause this operation was decomposed from");
1312 if (!getVar() || !getAccVar())
1313 return emitError("must have both host and device pointers");
1314 if (failed(checkVarAndVarType(*this)))
1315 return failure();
1316 if (failed(checkVarAndAccVar(*this)))
1317 return failure();
1318 if (failed(checkValidModifier(*this, acc::DataClauseModifier::zero |
1319 acc::DataClauseModifier::always |
1320 acc::DataClauseModifier::capture)))
1321 return failure();
1322 return success();
1323}
1324
1325bool acc::CopyoutOp::isCopyoutZero() {
1326 return getDataClause() == acc::DataClause::acc_copyout_zero ||
1327 acc::bitEnumContainsAny(getModifiers(), acc::DataClauseModifier::zero);
1328}
1329
1330//===----------------------------------------------------------------------===//
1331// DeleteOp
1332//===----------------------------------------------------------------------===//
1333LogicalResult acc::DeleteOp::verify() {
1334 // Test for all clauses this operation can be decomposed from:
1335 if (getDataClause() != acc::DataClause::acc_delete &&
1336 getDataClause() != acc::DataClause::acc_create &&
1337 getDataClause() != acc::DataClause::acc_create_zero &&
1338 getDataClause() != acc::DataClause::acc_copyin &&
1339 getDataClause() != acc::DataClause::acc_copyin_readonly &&
1340 getDataClause() != acc::DataClause::acc_present &&
1341 getDataClause() != acc::DataClause::acc_no_create &&
1342 getDataClause() != acc::DataClause::acc_declare_device_resident &&
1343 getDataClause() != acc::DataClause::acc_declare_link)
1344 return emitError(
1345 "data clause associated with delete operation must match its intent"
1346 " or specify original clause this operation was decomposed from");
1347 if (!getAccVar())
1348 return emitError("must have device pointer");
1349 // This op is the exit part of copyin and create - thus allow all modifiers
1350 // allowed on either case.
1351 if (failed(checkValidModifier(*this, acc::DataClauseModifier::zero |
1352 acc::DataClauseModifier::readonly |
1353 acc::DataClauseModifier::always |
1354 acc::DataClauseModifier::capture)))
1355 return failure();
1356 return success();
1357}
1358
1359//===----------------------------------------------------------------------===//
1360// DetachOp
1361//===----------------------------------------------------------------------===//
1362LogicalResult acc::DetachOp::verify() {
1363 // Test for all clauses this operation can be decomposed from:
1364 if (getDataClause() != acc::DataClause::acc_detach &&
1365 getDataClause() != acc::DataClause::acc_attach)
1366 return emitError(
1367 "data clause associated with detach operation must match its intent"
1368 " or specify original clause this operation was decomposed from");
1369 if (!getAccVar())
1370 return emitError("must have device pointer");
1371 if (failed(checkNoModifier(*this)))
1372 return failure();
1373 return success();
1374}
1375
1376//===----------------------------------------------------------------------===//
1377// HostOp
1378//===----------------------------------------------------------------------===//
1379LogicalResult acc::UpdateHostOp::verify() {
1380 // Test for all clauses this operation can be decomposed from:
1381 if (getDataClause() != acc::DataClause::acc_update_host &&
1382 getDataClause() != acc::DataClause::acc_update_self)
1383 return emitError(
1384 "data clause associated with host operation must match its intent"
1385 " or specify original clause this operation was decomposed from");
1386 if (!getVar() || !getAccVar())
1387 return emitError("must have both host and device pointers");
1388 if (failed(checkVarAndVarType(*this)))
1389 return failure();
1390 if (failed(checkVarAndAccVar(*this)))
1391 return failure();
1392 if (failed(checkNoModifier(*this)))
1393 return failure();
1394 return success();
1395}
1396
1397//===----------------------------------------------------------------------===//
1398// DeviceOp
1399//===----------------------------------------------------------------------===//
1400LogicalResult acc::UpdateDeviceOp::verify() {
1401 // Test for all clauses this operation can be decomposed from:
1402 if (getDataClause() != acc::DataClause::acc_update_device)
1403 return emitError(
1404 "data clause associated with device operation must match its intent"
1405 " or specify original clause this operation was decomposed from");
1406 if (failed(checkVarAndVarType(*this)))
1407 return failure();
1408 if (failed(checkVarAndAccVar(*this)))
1409 return failure();
1410 if (failed(checkNoModifier(*this)))
1411 return failure();
1412 return success();
1413}
1414
1415//===----------------------------------------------------------------------===//
1416// UseDeviceOp
1417//===----------------------------------------------------------------------===//
1418LogicalResult acc::UseDeviceOp::verify() {
1419 // Test for all clauses this operation can be decomposed from:
1420 if (getDataClause() != acc::DataClause::acc_use_device)
1421 return emitError(
1422 "data clause associated with use_device operation must match its intent"
1423 " or specify original clause this operation was decomposed from");
1424 if (failed(checkVarAndVarType(*this)))
1425 return failure();
1426 if (failed(checkVarAndAccVar(*this)))
1427 return failure();
1428 if (failed(checkNoModifier(*this)))
1429 return failure();
1430 return success();
1431}
1432
1433//===----------------------------------------------------------------------===//
1434// CacheOp
1435//===----------------------------------------------------------------------===//
1436LogicalResult acc::CacheOp::verify() {
1437 // Test for all clauses this operation can be decomposed from:
1438 if (getDataClause() != acc::DataClause::acc_cache &&
1439 getDataClause() != acc::DataClause::acc_cache_readonly)
1440 return emitError(
1441 "data clause associated with cache operation must match its intent"
1442 " or specify original clause this operation was decomposed from");
1443 if (failed(checkVarAndVarType(*this)))
1444 return failure();
1445 if (failed(checkVarAndAccVar(*this)))
1446 return failure();
1447 if (failed(checkValidModifier(*this, acc::DataClauseModifier::readonly)))
1448 return failure();
1449 return success();
1450}
1451
1452bool acc::CacheOp::isCacheReadonly() {
1453 return getDataClause() == acc::DataClause::acc_cache_readonly ||
1454 acc::bitEnumContainsAny(getModifiers(),
1455 acc::DataClauseModifier::readonly);
1456}
1457
1458//===----------------------------------------------------------------------===//
1459// Data entry/exit operations - getEffects implementations
1460//===----------------------------------------------------------------------===//
1461
1462// This function returns true iff the given operation is enclosed
1463// in any ACC_COMPUTE_CONSTRUCT_OPS operation.
1464// It is quite alike acc::getEnclosingComputeOp() utility,
1465// but we cannot use it here.
1469
1470/// Helper to add an effect on an operand, referenced by its mutable range.
1471template <typename EffectTy>
1474 &effects,
1475 MutableOperandRange operand) {
1476 for (unsigned i = 0, e = operand.size(); i < e; ++i)
1477 effects.emplace_back(EffectTy::get(), &operand[i]);
1478}
1479
1480/// Helper to add an effect on a result value.
1481template <typename EffectTy>
1484 &effects,
1485 Value result) {
1486 effects.emplace_back(EffectTy::get(), mlir::cast<mlir::OpResult>(result));
1487}
1488
1489// PrivateOp: accVar result write.
1490void acc::PrivateOp::getEffects(
1492 &effects) {
1493 // If acc.private is enclosed into a compute operation,
1494 // then it denotes the device side privatization, hence
1495 // it does not access the CurrentDeviceIdResource.
1496 if (!isEnclosedIntoComputeOp(getOperation()))
1497 effects.emplace_back(MemoryEffects::Read::get(),
1499 // TODO: should this be MemoryEffects::Allocate?
1501}
1502
1503// FirstprivateOp: var read, accVar result write.
1504void acc::FirstprivateOp::getEffects(
1506 &effects) {
1507 // If acc.firstprivate is enclosed into a compute operation,
1508 // then it denotes the device side privatization, hence
1509 // it does not access the CurrentDeviceIdResource.
1510 if (!isEnclosedIntoComputeOp(getOperation()))
1511 effects.emplace_back(MemoryEffects::Read::get(),
1513 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1515}
1516
1517// ReductionOp: var read, accVar result write.
1518void acc::ReductionOp::getEffects(
1520 &effects) {
1521 // If acc.reduction is enclosed into a compute operation,
1522 // then it denotes the device side reduction, hence
1523 // it does not access the CurrentDeviceIdResource.
1524 if (!isEnclosedIntoComputeOp(getOperation()))
1525 effects.emplace_back(MemoryEffects::Read::get(),
1527 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1529}
1530
1531// DevicePtrOp: RuntimeCounters read.
1532void acc::DevicePtrOp::getEffects(
1534 &effects) {
1535 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1536 effects.emplace_back(MemoryEffects::Read::get(),
1538}
1539
1540// PresentOp: RuntimeCounters read+write.
1541void acc::PresentOp::getEffects(
1543 &effects) {
1544 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1545 effects.emplace_back(MemoryEffects::Write::get(),
1547 effects.emplace_back(MemoryEffects::Read::get(),
1549}
1550
1551// CopyinOp: RuntimeCounters read+write, var read, accVar result write.
1552void acc::CopyinOp::getEffects(
1554 &effects) {
1555 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1556 effects.emplace_back(MemoryEffects::Write::get(),
1558 effects.emplace_back(MemoryEffects::Read::get(),
1560 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1562}
1563
1564// CreateOp: RuntimeCounters read+write, accVar result write.
1565void acc::CreateOp::getEffects(
1567 &effects) {
1568 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1569 effects.emplace_back(MemoryEffects::Write::get(),
1571 effects.emplace_back(MemoryEffects::Read::get(),
1573 // TODO: should this be MemoryEffects::Allocate?
1575}
1576
1577// NoCreateOp: RuntimeCounters read+write.
1578void acc::NoCreateOp::getEffects(
1580 &effects) {
1581 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1582 effects.emplace_back(MemoryEffects::Write::get(),
1584 effects.emplace_back(MemoryEffects::Read::get(),
1586}
1587
1588// AttachOp: RuntimeCounters read+write, var read.
1589void acc::AttachOp::getEffects(
1591 &effects) {
1592 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1593 effects.emplace_back(MemoryEffects::Write::get(),
1595 effects.emplace_back(MemoryEffects::Read::get(),
1597 // TODO: should we also add MemoryEffects::Write?
1598 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1599}
1600
1601// GetDevicePtrOp: RuntimeCounters read.
1602void acc::GetDevicePtrOp::getEffects(
1604 &effects) {
1605 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1606 effects.emplace_back(MemoryEffects::Read::get(),
1608}
1609
1610// UpdateDeviceOp: var read, accVar result write.
1611void acc::UpdateDeviceOp::getEffects(
1613 &effects) {
1614 effects.emplace_back(MemoryEffects::Read::get(),
1616 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1618}
1619
1620// UseDeviceOp: RuntimeCounters read.
1621void acc::UseDeviceOp::getEffects(
1623 &effects) {
1624 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1625 effects.emplace_back(MemoryEffects::Read::get(),
1627}
1628
1629// DeclareDeviceResidentOp: RuntimeCounters write, var read.
1630void acc::DeclareDeviceResidentOp::getEffects(
1632 &effects) {
1633 effects.emplace_back(MemoryEffects::Write::get(),
1635 effects.emplace_back(MemoryEffects::Read::get(),
1637 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1638}
1639
1640// DeclareLinkOp: RuntimeCounters write, var read.
1641void acc::DeclareLinkOp::getEffects(
1643 &effects) {
1644 effects.emplace_back(MemoryEffects::Write::get(),
1646 effects.emplace_back(MemoryEffects::Read::get(),
1648 addOperandEffect<MemoryEffects::Read>(effects, getVarMutable());
1649}
1650
1651// CacheOp: NoMemoryEffect
1652void acc::CacheOp::getEffects(
1654 &effects) {}
1655
1656// CopyoutOp: RuntimeCounters read+write, accVar read, var write.
1657void acc::CopyoutOp::getEffects(
1659 &effects) {
1660 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1661 effects.emplace_back(MemoryEffects::Write::get(),
1663 effects.emplace_back(MemoryEffects::Read::get(),
1665 addOperandEffect<MemoryEffects::Read>(effects, getAccVarMutable());
1666 addOperandEffect<MemoryEffects::Write>(effects, getVarMutable());
1667}
1668
1669// DeleteOp: RuntimeCounters read+write, accVar read.
1670void acc::DeleteOp::getEffects(
1672 &effects) {
1673 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1674 effects.emplace_back(MemoryEffects::Write::get(),
1676 effects.emplace_back(MemoryEffects::Read::get(),
1678 addOperandEffect<MemoryEffects::Read>(effects, getAccVarMutable());
1679}
1680
1681// DetachOp: RuntimeCounters read+write, accVar read.
1682void acc::DetachOp::getEffects(
1684 &effects) {
1685 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1686 effects.emplace_back(MemoryEffects::Write::get(),
1688 effects.emplace_back(MemoryEffects::Read::get(),
1690 addOperandEffect<MemoryEffects::Read>(effects, getAccVarMutable());
1691}
1692
1693// UpdateHostOp: RuntimeCounters read+write, accVar read, var write.
1694void acc::UpdateHostOp::getEffects(
1696 &effects) {
1697 effects.emplace_back(MemoryEffects::Read::get(), acc::RuntimeCounters::get());
1698 effects.emplace_back(MemoryEffects::Write::get(),
1700 effects.emplace_back(MemoryEffects::Read::get(),
1702 addOperandEffect<MemoryEffects::Read>(effects, getAccVarMutable());
1703 addOperandEffect<MemoryEffects::Write>(effects, getVarMutable());
1704}
1705
1706namespace {
1707/// Pattern to remove operation without region that have constant false `ifCond`
1708/// and remove the condition from the operation if the `ifCond` is a true
1709/// constant.
1710template <typename OpTy>
1711struct RemoveConstantIfCondition : public OpRewritePattern<OpTy> {
1712 using OpRewritePattern<OpTy>::OpRewritePattern;
1713
1714 LogicalResult matchAndRewrite(OpTy op,
1715 PatternRewriter &rewriter) const override {
1716 // Early return if there is no condition.
1717 Value ifCond = op.getIfCond();
1718 if (!ifCond)
1719 return failure();
1720
1721 IntegerAttr constAttr;
1722 if (!matchPattern(ifCond, m_Constant(&constAttr)))
1723 return failure();
1724 if (constAttr.getInt())
1725 rewriter.modifyOpInPlace(op, [&]() { op.getIfCondMutable().erase(0); });
1726 else
1727 rewriter.eraseOp(op);
1728
1729 return success();
1730 }
1731};
1732
1733/// Replaces the given op with the contents of the given single-block region,
1734/// using the operands of the block terminator to replace operation results.
1735static void replaceOpWithRegion(PatternRewriter &rewriter, Operation *op,
1736 Region &region, ValueRange blockArgs = {}) {
1737 assert(region.hasOneBlock() && "expected single-block region");
1738 Block *block = &region.front();
1739 Operation *terminator = block->getTerminator();
1740 ValueRange results = terminator->getOperands();
1741 rewriter.inlineBlockBefore(block, op, blockArgs);
1742 rewriter.replaceOp(op, results);
1743 rewriter.eraseOp(terminator);
1744}
1745
1746/// Pattern to remove operation with region that have constant false `ifCond`
1747/// and remove the condition from the operation if the `ifCond` is constant
1748/// true.
1749template <typename OpTy>
1750struct RemoveConstantIfConditionWithRegion : public OpRewritePattern<OpTy> {
1751 using OpRewritePattern<OpTy>::OpRewritePattern;
1752
1753 LogicalResult matchAndRewrite(OpTy op,
1754 PatternRewriter &rewriter) const override {
1755 // Early return if there is no condition.
1756 Value ifCond = op.getIfCond();
1757 if (!ifCond)
1758 return failure();
1759
1760 IntegerAttr constAttr;
1761 if (!matchPattern(ifCond, m_Constant(&constAttr)))
1762 return failure();
1763 if (constAttr.getInt())
1764 rewriter.modifyOpInPlace(op, [&]() { op.getIfCondMutable().erase(0); });
1765 else
1766 replaceOpWithRegion(rewriter, op, op.getRegion());
1767
1768 return success();
1769 }
1770};
1771
1772//===----------------------------------------------------------------------===//
1773// Recipe Region Helpers
1774//===----------------------------------------------------------------------===//
1775
1776/// Create and populate an init region for privatization recipes.
1777/// Returns success if the region is populated, failure otherwise.
1778/// Sets needsFree to indicate if the allocated memory requires deallocation.
1779/// The `hostVar` is the original host variable used to derive
1780/// language-specific metadata via `genPrivateVariableInfo`.
1781/// The `varInfo` output parameter is set to the variable info produced.
1782static LogicalResult createInitRegion(OpBuilder &builder, Location loc,
1783 Region &initRegion, Value hostVar,
1784 StringRef varName, ValueRange bounds,
1785 bool &needsFree,
1786 acc::VariableInfoAttr &varInfo,
1787 SmallVectorImpl<Value> &destroyValues) {
1788 Type varType = hostVar.getType();
1789
1790 // Create init block with arguments: original value + bounds
1791 SmallVector<Type> argTypes{varType};
1792 SmallVector<Location> argLocs{loc};
1793 for (Value bound : bounds) {
1794 argTypes.push_back(bound.getType());
1795 argLocs.push_back(loc);
1796 }
1797
1798 Block *initBlock = builder.createBlock(&initRegion);
1799 initBlock->addArguments(argTypes, argLocs);
1800 builder.setInsertionPointToStart(initBlock);
1801
1802 Value privatizedValue;
1803
1804 // Get the block argument that represents the original variable
1805 Value blockArgVar = initBlock->getArgument(0);
1806
1807 // Generate init region body based on variable type
1808 if (isa<MappableType>(varType)) {
1809 auto mappableTy = cast<MappableType>(varType);
1810 auto typedVar = cast<TypedValue<MappableType>>(blockArgVar);
1811 auto typedHostVar = cast<TypedValue<MappableType>>(hostVar);
1812 varInfo = mappableTy.genPrivateVariableInfo(typedHostVar);
1813 privatizedValue =
1814 mappableTy.generatePrivateInit(builder, loc, typedVar, varName, bounds,
1815 {}, varInfo, needsFree, destroyValues);
1816 if (!privatizedValue)
1817 return failure();
1818 } else {
1819 assert(isa<PointerLikeType>(varType) && "Expected PointerLikeType");
1820 auto pointerLikeTy = cast<PointerLikeType>(varType);
1821 // Use PointerLikeType's allocation API with the block argument
1822 privatizedValue = pointerLikeTy.genAllocate(builder, loc, varName, varType,
1823 blockArgVar, needsFree);
1824 if (!privatizedValue)
1825 return failure();
1826 }
1827
1828 // Add yield operation to init block
1829 SmallVector<Value> initResults{privatizedValue};
1830 initResults.append(destroyValues);
1831 acc::YieldOp::create(builder, loc, initResults);
1832
1833 return success();
1834}
1835
1836/// Create and populate a copy region for firstprivate recipes.
1837/// Returns success if the region is populated, failure otherwise.
1838/// `varInfo` must be the attribute produced by `createInitRegion` for
1839/// `MappableType` (it is unused for `PointerLikeType` copy paths).
1840static LogicalResult createCopyRegion(OpBuilder &builder, Location loc,
1841 Region &copyRegion, Type varType,
1842 ValueRange bounds,
1843 acc::VariableInfoAttr varInfo) {
1844 // Create copy block with arguments: original value + privatized value +
1845 // bounds
1846 SmallVector<Type> copyArgTypes{varType, varType};
1847 SmallVector<Location> copyArgLocs{loc, loc};
1848 for (Value bound : bounds) {
1849 copyArgTypes.push_back(bound.getType());
1850 copyArgLocs.push_back(loc);
1851 }
1852
1853 Block *copyBlock = builder.createBlock(&copyRegion);
1854 copyBlock->addArguments(copyArgTypes, copyArgLocs);
1855 builder.setInsertionPointToStart(copyBlock);
1856
1857 Value originalArg = copyBlock->getArgument(0);
1858 Value privatizedArg = copyBlock->getArgument(1);
1859
1860 if (isa<MappableType>(varType)) {
1861 auto mappableTy = cast<MappableType>(varType);
1862 // generateCopy(src, dest): copy from original (arg0) into privatized
1863 // (arg1).
1864 if (!mappableTy.generateCopy(
1865 builder, loc, cast<TypedValue<MappableType>>(originalArg),
1866 cast<TypedValue<MappableType>>(privatizedArg), bounds, varInfo))
1867 return failure();
1868 } else {
1869 assert(isa<PointerLikeType>(varType) && "Expected PointerLikeType");
1870 auto pointerLikeTy = cast<PointerLikeType>(varType);
1871 if (!pointerLikeTy.genCopy(
1872 builder, loc, cast<TypedValue<PointerLikeType>>(privatizedArg),
1873 cast<TypedValue<PointerLikeType>>(originalArg), varType))
1874 return failure();
1875 }
1876
1877 // Add terminator to copy block
1878 acc::TerminatorOp::create(builder, loc);
1879
1880 return success();
1881}
1882
1883/// Create and populate a destroy region for privatization recipes.
1884/// Returns success if the region is populated, failure otherwise.
1885/// The `varInfo` carries language-specific metadata produced by
1886/// `createInitRegion`.
1887static LogicalResult
1888createDestroyRegion(OpBuilder &builder, Location loc, Region &destroyRegion,
1889 Type varType, Value allocRes, ValueRange destroyValues,
1890 ValueRange bounds, acc::VariableInfoAttr varInfo) {
1891 // Create destroy block with arguments: original value + privatized value +
1892 // values preserved for destruction + bounds.
1893 SmallVector<Type> destroyArgTypes{varType, varType};
1894 SmallVector<Location> destroyArgLocs{loc, loc};
1895 for (Value destroyValue : destroyValues) {
1896 destroyArgTypes.push_back(destroyValue.getType());
1897 destroyArgLocs.push_back(loc);
1898 }
1899 for (Value bound : bounds) {
1900 destroyArgTypes.push_back(bound.getType());
1901 destroyArgLocs.push_back(loc);
1902 }
1903
1904 Block *destroyBlock = builder.createBlock(&destroyRegion);
1905 destroyBlock->addArguments(destroyArgTypes, destroyArgLocs);
1906 builder.setInsertionPointToStart(destroyBlock);
1907
1908 auto varToFree =
1909 cast<TypedValue<PointerLikeType>>(destroyBlock->getArgument(1));
1910 if (isa<MappableType>(varType)) {
1911 auto mappableTy = cast<MappableType>(varType);
1912 ValueRange destroyArgs =
1913 destroyBlock->getArguments().slice(2, destroyValues.size());
1914 ValueRange destroyBounds =
1915 destroyBlock->getArguments().drop_front(2 + destroyValues.size());
1916 if (!mappableTy.generatePrivateDestroy(builder, loc, varToFree, destroyArgs,
1917 destroyBounds, varInfo))
1918 return failure();
1919 } else {
1920 assert(isa<PointerLikeType>(varType) && "Expected PointerLikeType");
1921 auto pointerLikeTy = cast<PointerLikeType>(varType);
1922 if (!pointerLikeTy.genFree(builder, loc, varToFree, allocRes, varType))
1923 return failure();
1924 }
1925
1926 acc::TerminatorOp::create(builder, loc);
1927 return success();
1928}
1929
1930} // namespace
1931
1932//===----------------------------------------------------------------------===//
1933// PrivateRecipeOp
1934//===----------------------------------------------------------------------===//
1935
1937 Operation *op, Region &region, StringRef regionType, StringRef regionName,
1938 Type type, bool verifyYield, bool optional = false) {
1939 if (optional && region.empty())
1940 return success();
1941
1942 if (region.empty())
1943 return op->emitOpError() << "expects non-empty " << regionName << " region";
1944 Block &firstBlock = region.front();
1945 if (firstBlock.getNumArguments() < 1 ||
1946 firstBlock.getArgument(0).getType() != type)
1947 return op->emitOpError() << "expects " << regionName
1948 << " region first "
1949 "argument of the "
1950 << regionType << " type";
1951
1952 if (verifyYield) {
1953 for (YieldOp yieldOp : region.getOps<acc::YieldOp>()) {
1954 if (yieldOp.getOperands().size() != 1 ||
1955 yieldOp.getOperands().getTypes()[0] != type)
1956 return op->emitOpError() << "expects " << regionName
1957 << " region to "
1958 "yield a value of the "
1959 << regionType << " type";
1960 }
1961 }
1962 return success();
1963}
1964
1965LogicalResult acc::PrivateRecipeOp::verifyRegions() {
1966 if (failed(verifyInitLikeSingleArgRegion(*this, getInitRegion(),
1967 "privatization", "init", getType(),
1968 /*verifyYield=*/false)))
1969 return failure();
1971 *this, getDestroyRegion(), "privatization", "destroy", getType(),
1972 /*verifyYield=*/false, /*optional=*/true)))
1973 return failure();
1974 return success();
1975}
1976
1977std::optional<PrivateRecipeOp>
1978PrivateRecipeOp::createAndPopulate(OpBuilder &builder, Location loc,
1979 StringRef recipeName, Value hostVar,
1980 StringRef varName, ValueRange bounds) {
1981 Type varType = hostVar.getType();
1982
1983 // First, validate that we can handle this variable type
1984 bool isMappable = isa<MappableType>(varType);
1985 bool isPointerLike = isa<PointerLikeType>(varType);
1986
1987 // Unsupported type
1988 if (!isMappable && !isPointerLike)
1989 return std::nullopt;
1990
1991 OpBuilder::InsertionGuard guard(builder);
1992
1993 // Create the recipe operation first so regions have proper parent context
1994 auto recipe = PrivateRecipeOp::create(builder, loc, recipeName,
1995 /*sym_visibility=*/nullptr, varType);
1996
1997 // Populate the init region
1998 bool needsFree = false;
1999 acc::VariableInfoAttr varInfo;
2000 SmallVector<Value> destroyValues;
2001 if (failed(createInitRegion(builder, loc, recipe.getInitRegion(), hostVar,
2002 varName, bounds, needsFree, varInfo,
2003 destroyValues))) {
2004 recipe.erase();
2005 return std::nullopt;
2006 }
2007
2008 // Only create destroy region if the allocation needs deallocation
2009 if (needsFree) {
2010 // Extract the allocated value from the init block's yield operation
2011 auto yieldOp =
2012 cast<acc::YieldOp>(recipe.getInitRegion().front().getTerminator());
2013 Value allocRes = yieldOp.getOperand(0);
2014
2015 if (failed(createDestroyRegion(builder, loc, recipe.getDestroyRegion(),
2016 varType, allocRes, destroyValues, bounds,
2017 varInfo))) {
2018 recipe.erase();
2019 return std::nullopt;
2020 }
2021 }
2022
2023 return recipe;
2024}
2025
2026std::optional<PrivateRecipeOp>
2027PrivateRecipeOp::createAndPopulate(OpBuilder &builder, Location loc,
2028 StringRef recipeName,
2029 FirstprivateRecipeOp firstprivRecipe) {
2030 // Create the private.recipe op with the same type as the firstprivate.recipe.
2031 OpBuilder::InsertionGuard guard(builder);
2032 auto varType = firstprivRecipe.getType();
2033 auto recipe = PrivateRecipeOp::create(builder, loc, recipeName,
2034 /*sym_visibility=*/nullptr, varType);
2035
2036 // Clone the init region
2037 IRMapping mapping;
2038 firstprivRecipe.getInitRegion().cloneInto(&recipe.getInitRegion(), mapping);
2039
2040 // Clone destroy region if the firstprivate.recipe has one.
2041 if (!firstprivRecipe.getDestroyRegion().empty()) {
2042 IRMapping mapping;
2043 firstprivRecipe.getDestroyRegion().cloneInto(&recipe.getDestroyRegion(),
2044 mapping);
2045 }
2046 return recipe;
2047}
2048
2049//===----------------------------------------------------------------------===//
2050// FirstprivateRecipeOp
2051//===----------------------------------------------------------------------===//
2052
2053LogicalResult acc::FirstprivateRecipeOp::verifyRegions() {
2054 if (failed(verifyInitLikeSingleArgRegion(*this, getInitRegion(),
2055 "privatization", "init", getType(),
2056 /*verifyYield=*/false)))
2057 return failure();
2058
2059 if (getCopyRegion().empty())
2060 return emitOpError() << "expects non-empty copy region";
2061
2062 Block &firstBlock = getCopyRegion().front();
2063 if (firstBlock.getNumArguments() < 2 ||
2064 firstBlock.getArgument(0).getType() != getType())
2065 return emitOpError() << "expects copy region with two arguments of the "
2066 "privatization type";
2067
2068 if (getDestroyRegion().empty())
2069 return success();
2070
2071 if (failed(verifyInitLikeSingleArgRegion(*this, getDestroyRegion(),
2072 "privatization", "destroy",
2073 getType(), /*verifyYield=*/false)))
2074 return failure();
2075
2076 return success();
2077}
2078
2079std::optional<FirstprivateRecipeOp>
2080FirstprivateRecipeOp::createAndPopulate(OpBuilder &builder, Location loc,
2081 StringRef recipeName, Value hostVar,
2082 StringRef varName, ValueRange bounds) {
2083 Type varType = hostVar.getType();
2084
2085 // First, validate that we can handle this variable type
2086 bool isMappable = isa<MappableType>(varType);
2087 bool isPointerLike = isa<PointerLikeType>(varType);
2088
2089 // Unsupported type
2090 if (!isMappable && !isPointerLike)
2091 return std::nullopt;
2092
2093 OpBuilder::InsertionGuard guard(builder);
2094
2095 // Create the recipe operation first so regions have proper parent context
2096 auto recipe = FirstprivateRecipeOp::create(
2097 builder, loc, recipeName, /*sym_visibility=*/nullptr, varType);
2098
2099 // Populate the init region
2100 bool needsFree = false;
2101 // Filled by createInitRegion for mappable variables (genPrivateVariableInfo);
2102 // then passed through to copy/destroy so generateCopy /
2103 // generatePrivateDestroy receive the same metadata as generatePrivateInit.
2104 acc::VariableInfoAttr varInfo;
2105 SmallVector<Value> destroyValues;
2106 if (failed(createInitRegion(builder, loc, recipe.getInitRegion(), hostVar,
2107 varName, bounds, needsFree, varInfo,
2108 destroyValues))) {
2109 recipe.erase();
2110 return std::nullopt;
2111 }
2112
2113 // Populate the copy region (uses varInfo for MappableType::generateCopy).
2114 if (failed(createCopyRegion(builder, loc, recipe.getCopyRegion(), varType,
2115 bounds, varInfo))) {
2116 recipe.erase();
2117 return std::nullopt;
2118 }
2119
2120 // Only create destroy region if the allocation needs deallocation
2121 if (needsFree) {
2122 // Extract the allocated value from the init block's yield operation
2123 auto yieldOp =
2124 cast<acc::YieldOp>(recipe.getInitRegion().front().getTerminator());
2125 Value allocRes = yieldOp.getOperand(0);
2126
2127 if (failed(createDestroyRegion(builder, loc, recipe.getDestroyRegion(),
2128 varType, allocRes, destroyValues, bounds,
2129 varInfo))) {
2130 recipe.erase();
2131 return std::nullopt;
2132 }
2133 }
2134
2135 return recipe;
2136}
2137
2138//===----------------------------------------------------------------------===//
2139// ReductionRecipeOp
2140//===----------------------------------------------------------------------===//
2141
2142LogicalResult acc::ReductionRecipeOp::verifyRegions() {
2143 if (failed(verifyInitLikeSingleArgRegion(*this, getInitRegion(), "reduction",
2144 "init", getType(),
2145 /*verifyYield=*/false)))
2146 return failure();
2147
2148 if (getCombinerRegion().empty())
2149 return emitOpError() << "expects non-empty combiner region";
2150
2151 Block &reductionBlock = getCombinerRegion().front();
2152 if (reductionBlock.getNumArguments() < 2 ||
2153 reductionBlock.getArgument(0).getType() != getType() ||
2154 reductionBlock.getArgument(1).getType() != getType())
2155 return emitOpError() << "expects combiner region with the first two "
2156 << "arguments of the reduction type";
2157
2158 for (YieldOp yieldOp : getCombinerRegion().getOps<YieldOp>()) {
2159 if (yieldOp.getOperands().size() != 1 ||
2160 yieldOp.getOperands().getTypes()[0] != getType())
2161 return emitOpError() << "expects combiner region to yield a value "
2162 "of the reduction type";
2163 }
2164
2165 return success();
2166}
2167
2168//===----------------------------------------------------------------------===//
2169// ParallelOp
2170//===----------------------------------------------------------------------===//
2171
2172/// Check dataOperands for acc.parallel, acc.serial and acc.kernels.
2173template <typename Op>
2174static LogicalResult checkDataOperands(Op op,
2175 const mlir::ValueRange &operands) {
2176 for (mlir::Value operand : operands)
2177 if (!mlir::isa<acc::AttachOp, acc::CopyinOp, acc::CopyoutOp, acc::CreateOp,
2178 acc::DeleteOp, acc::DetachOp, acc::DevicePtrOp,
2179 acc::GetDevicePtrOp, acc::NoCreateOp, acc::PresentOp,
2180 acc::MapInfoOp>(operand.getDefiningOp()))
2181 return op.emitError(
2182 "expect data entry/exit operation or acc.getdeviceptr "
2183 "as defining op");
2184 return success();
2185}
2186
2187template <typename OpT, typename RecipeOpT>
2188static LogicalResult checkPrivateOperands(mlir::Operation *accConstructOp,
2189 const mlir::ValueRange &operands,
2190 llvm::StringRef operandName) {
2192 for (mlir::Value operand : operands) {
2193 if (!mlir::isa<OpT>(operand.getDefiningOp()))
2194 return accConstructOp->emitOpError()
2195 << "expected " << operandName << " as defining op";
2196 if (!set.insert(operand).second)
2197 return accConstructOp->emitOpError()
2198 << operandName << " operand appears more than once";
2199 }
2200 return success();
2201}
2202
2203unsigned ParallelOp::getNumDataOperands() {
2204 return getReductionOperands().size() + getPrivateOperands().size() +
2205 getFirstprivateOperands().size() + getDataClauseOperands().size();
2206}
2207
2208Value ParallelOp::getDataOperand(unsigned i) {
2209 unsigned numOptional = getAsyncOperands().size();
2210 numOptional += getNumGangs().size();
2211 numOptional += getNumWorkers().size();
2212 numOptional += getVectorLength().size();
2213 numOptional += getIfCond() ? 1 : 0;
2214 numOptional += getSelfCond() ? 1 : 0;
2215 return getOperand(getWaitOperands().size() + numOptional + i);
2216}
2217
2218template <typename Op>
2219static LogicalResult verifyDeviceTypeCountMatch(Op op, OperandRange operands,
2220 ArrayAttr deviceTypes,
2221 llvm::StringRef keyword) {
2222 if (!operands.empty() &&
2223 (!deviceTypes || deviceTypes.getValue().size() != operands.size()))
2224 return op.emitOpError() << keyword << " operands count must match "
2225 << keyword << " device_type count";
2226 return success();
2227}
2228
2229template <typename Op>
2231 Op op, OperandRange operands, DenseI32ArrayAttr segments,
2232 ArrayAttr deviceTypes, llvm::StringRef keyword, int32_t maxInSegment = 0) {
2233 std::size_t numOperandsInSegments = 0;
2234 std::size_t nbOfSegments = 0;
2235
2236 if (segments) {
2237 for (auto segCount : segments.asArrayRef()) {
2238 if (maxInSegment != 0 && segCount > maxInSegment)
2239 return op.emitOpError() << keyword << " expects a maximum of "
2240 << maxInSegment << " values per segment";
2241 numOperandsInSegments += segCount;
2242 ++nbOfSegments;
2243 }
2244 }
2245
2246 if ((numOperandsInSegments != operands.size()) ||
2247 (!deviceTypes && !operands.empty()))
2248 return op.emitOpError()
2249 << keyword << " operand count does not match count in segments";
2250 if (deviceTypes && deviceTypes.getValue().size() != nbOfSegments)
2251 return op.emitOpError()
2252 << keyword << " segment count does not match device_type count";
2253 return success();
2254}
2255
2256LogicalResult acc::ParallelOp::verify() {
2257 if (failed(checkPrivateOperands<mlir::acc::PrivateOp,
2258 mlir::acc::PrivateRecipeOp>(
2259 *this, getPrivateOperands(), "private")))
2260 return failure();
2261 if (failed(checkPrivateOperands<mlir::acc::FirstprivateOp,
2262 mlir::acc::FirstprivateRecipeOp>(
2263 *this, getFirstprivateOperands(), "firstprivate")))
2264 return failure();
2265 if (failed(checkPrivateOperands<mlir::acc::ReductionOp,
2266 mlir::acc::ReductionRecipeOp>(
2267 *this, getReductionOperands(), "reduction")))
2268 return failure();
2269
2271 *this, getNumGangs(), getNumGangsSegmentsAttr(),
2272 getNumGangsDeviceTypeAttr(), "num_gangs", 3)))
2273 return failure();
2274
2276 *this, getWaitOperands(), getWaitOperandsSegmentsAttr(),
2277 getWaitOperandsDeviceTypeAttr(), "wait")))
2278 return failure();
2279
2280 if (failed(verifyDeviceTypeCountMatch(*this, getNumWorkers(),
2281 getNumWorkersDeviceTypeAttr(),
2282 "num_workers")))
2283 return failure();
2284
2285 if (failed(verifyDeviceTypeCountMatch(*this, getVectorLength(),
2286 getVectorLengthDeviceTypeAttr(),
2287 "vector_length")))
2288 return failure();
2289
2291 getAsyncOperandsDeviceTypeAttr(),
2292 "async")))
2293 return failure();
2294
2296 return failure();
2297
2298 return checkDataOperands<acc::ParallelOp>(*this, getDataClauseOperands());
2299}
2300
2301static mlir::Value
2302getValueInDeviceTypeSegment(std::optional<mlir::ArrayAttr> arrayAttr,
2304 mlir::acc::DeviceType deviceType) {
2305 if (!arrayAttr)
2306 return {};
2307 if (auto pos = findSegment(*arrayAttr, deviceType))
2308 return range[*pos];
2309 return {};
2310}
2311
2312bool acc::ParallelOp::hasAsyncOnly() {
2313 return hasAsyncOnly(mlir::acc::DeviceType::None);
2314}
2315
2316bool acc::ParallelOp::hasAsyncOnly(mlir::acc::DeviceType deviceType) {
2317 return hasDeviceType(getAsyncOnly(), deviceType);
2318}
2319
2320mlir::Value acc::ParallelOp::getAsyncValue() {
2321 return getAsyncValue(mlir::acc::DeviceType::None);
2322}
2323
2324mlir::Value acc::ParallelOp::getAsyncValue(mlir::acc::DeviceType deviceType) {
2326 getAsyncOperands(), deviceType);
2327}
2328
2329mlir::Value acc::ParallelOp::getNumWorkersValue() {
2330 return getNumWorkersValue(mlir::acc::DeviceType::None);
2331}
2332
2334acc::ParallelOp::getNumWorkersValue(mlir::acc::DeviceType deviceType) {
2335 return getValueInDeviceTypeSegment(getNumWorkersDeviceType(), getNumWorkers(),
2336 deviceType);
2337}
2338
2339mlir::Value acc::ParallelOp::getVectorLengthValue() {
2340 return getVectorLengthValue(mlir::acc::DeviceType::None);
2341}
2342
2344acc::ParallelOp::getVectorLengthValue(mlir::acc::DeviceType deviceType) {
2345 return getValueInDeviceTypeSegment(getVectorLengthDeviceType(),
2346 getVectorLength(), deviceType);
2347}
2348
2349mlir::Operation::operand_range ParallelOp::getNumGangsValues() {
2350 return getNumGangsValues(mlir::acc::DeviceType::None);
2351}
2352
2354ParallelOp::getNumGangsValues(mlir::acc::DeviceType deviceType) {
2355 return getValuesFromSegments(getNumGangsDeviceType(), getNumGangs(),
2356 getNumGangsSegments(), deviceType);
2357}
2358
2360 std::optional<mlir::ArrayAttr> numGangsDeviceType,
2362 std::optional<llvm::ArrayRef<int32_t>> numGangsSegments,
2363 std::optional<mlir::ArrayAttr> numWorkersDeviceType,
2365 std::optional<mlir::ArrayAttr> vectorLengthDeviceType,
2366 mlir::Operation::operand_range vectorLength,
2367 mlir::acc::DeviceType deviceType) {
2368 return !getValuesFromSegments(numGangsDeviceType, numGangs, numGangsSegments,
2369 deviceType)
2370 .empty() ||
2371 getValueInDeviceTypeSegment(numWorkersDeviceType, numWorkers,
2372 deviceType) ||
2373 getValueInDeviceTypeSegment(vectorLengthDeviceType, vectorLength,
2374 deviceType);
2375}
2376
2377bool acc::ParallelOp::hasAnyGangWorkerVector(mlir::acc::DeviceType deviceType) {
2379 getNumGangsDeviceType(), getNumGangs(), getNumGangsSegments(),
2380 getNumWorkersDeviceType(), getNumWorkers(), getVectorLengthDeviceType(),
2381 getVectorLength(), deviceType);
2382}
2383
2384bool acc::ParallelOp::isEffectivelySerial() {
2385 return isGangWorkerVectorAllOne(*this);
2386}
2387
2388bool acc::ParallelOp::hasWaitOnly() {
2389 return hasWaitOnly(mlir::acc::DeviceType::None);
2390}
2391
2392bool acc::ParallelOp::hasWaitOnly(mlir::acc::DeviceType deviceType) {
2393 return hasDeviceType(getWaitOnly(), deviceType);
2394}
2395
2396mlir::Operation::operand_range ParallelOp::getWaitValues() {
2397 return getWaitValues(mlir::acc::DeviceType::None);
2398}
2399
2401ParallelOp::getWaitValues(mlir::acc::DeviceType deviceType) {
2403 getWaitOperandsDeviceType(), getWaitOperands(), getWaitOperandsSegments(),
2404 getHasWaitDevnum(), deviceType);
2405}
2406
2407mlir::Value ParallelOp::getWaitDevnum() {
2408 return getWaitDevnum(mlir::acc::DeviceType::None);
2409}
2410
2411mlir::Value ParallelOp::getWaitDevnum(mlir::acc::DeviceType deviceType) {
2412 return getWaitDevnumValue(getWaitOperandsDeviceType(), getWaitOperands(),
2413 getWaitOperandsSegments(), getHasWaitDevnum(),
2414 deviceType);
2415}
2416
2417void ParallelOp::build(mlir::OpBuilder &odsBuilder,
2418 mlir::OperationState &odsState,
2419 mlir::ValueRange numGangs, mlir::ValueRange numWorkers,
2420 mlir::ValueRange vectorLength,
2421 mlir::ValueRange asyncOperands,
2422 mlir::ValueRange waitOperands, mlir::Value ifCond,
2423 mlir::Value selfCond, mlir::ValueRange reductionOperands,
2424 mlir::ValueRange gangPrivateOperands,
2425 mlir::ValueRange gangFirstPrivateOperands,
2426 mlir::ValueRange dataClauseOperands) {
2427 ParallelOp::build(
2428 odsBuilder, odsState, asyncOperands, /*asyncOperandsDeviceType=*/nullptr,
2429 /*asyncOnly=*/nullptr, waitOperands, /*waitOperandsSegments=*/nullptr,
2430 /*waitOperandsDeviceType=*/nullptr, /*hasWaitDevnum=*/nullptr,
2431 /*waitOnly=*/nullptr, numGangs, /*numGangsSegments=*/nullptr,
2432 /*numGangsDeviceType=*/nullptr, numWorkers,
2433 /*numWorkersDeviceType=*/nullptr, vectorLength,
2434 /*vectorLengthDeviceType=*/nullptr, ifCond, selfCond,
2435 /*selfAttr=*/false, reductionOperands, gangPrivateOperands,
2436 gangFirstPrivateOperands, dataClauseOperands,
2437 /*defaultAttr=*/nullptr, /*combined=*/false);
2438}
2439
2440void acc::ParallelOp::addNumWorkersOperand(
2441 MLIRContext *context, mlir::Value newValue,
2442 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2443 setNumWorkersDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
2444 context, getNumWorkersDeviceTypeAttr(), effectiveDeviceTypes, newValue,
2445 getNumWorkersMutable()));
2446}
2447void acc::ParallelOp::addVectorLengthOperand(
2448 MLIRContext *context, mlir::Value newValue,
2449 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2450 setVectorLengthDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
2451 context, getVectorLengthDeviceTypeAttr(), effectiveDeviceTypes, newValue,
2452 getVectorLengthMutable()));
2453}
2454
2455void acc::ParallelOp::addAsyncOnly(
2456 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2457 setAsyncOnlyAttr(addDeviceTypeAffectedOperandHelper(
2458 context, getAsyncOnlyAttr(), effectiveDeviceTypes));
2459}
2460
2461void acc::ParallelOp::addAsyncOperand(
2462 MLIRContext *context, mlir::Value newValue,
2463 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2464 setAsyncOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
2465 context, getAsyncOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValue,
2466 getAsyncOperandsMutable()));
2467}
2468
2469void acc::ParallelOp::addNumGangsOperands(
2470 MLIRContext *context, mlir::ValueRange newValues,
2471 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2473 if (getNumGangsSegments())
2474 llvm::copy(*getNumGangsSegments(), std::back_inserter(segments));
2475
2476 setNumGangsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
2477 context, getNumGangsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
2478 getNumGangsMutable(), segments));
2479
2480 setNumGangsSegments(segments);
2481}
2482void acc::ParallelOp::addWaitOnly(
2483 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2484 setWaitOnlyAttr(addDeviceTypeAffectedOperandHelper(context, getWaitOnlyAttr(),
2485 effectiveDeviceTypes));
2486}
2487void acc::ParallelOp::addWaitOperands(
2488 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
2489 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
2490
2492 if (getWaitOperandsSegments())
2493 llvm::copy(*getWaitOperandsSegments(), std::back_inserter(segments));
2494
2495 setWaitOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
2496 context, getWaitOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
2497 getWaitOperandsMutable(), segments));
2498 setWaitOperandsSegments(segments);
2499
2501 if (getHasWaitDevnumAttr())
2502 llvm::copy(getHasWaitDevnumAttr(), std::back_inserter(hasDevnums));
2503 hasDevnums.insert(
2504 hasDevnums.end(),
2505 std::max(effectiveDeviceTypes.size(), static_cast<size_t>(1)),
2506 mlir::BoolAttr::get(context, hasDevnum));
2507 setHasWaitDevnumAttr(mlir::ArrayAttr::get(context, hasDevnums));
2508}
2509
2510void acc::ParallelOp::addPrivatization(MLIRContext *context,
2511 mlir::acc::PrivateOp op,
2512 mlir::acc::PrivateRecipeOp recipe) {
2513 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
2514 getPrivateOperandsMutable().append(op.getResult());
2515}
2516
2517void acc::ParallelOp::addFirstPrivatization(
2518 MLIRContext *context, mlir::acc::FirstprivateOp op,
2519 mlir::acc::FirstprivateRecipeOp recipe) {
2520 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
2521 getFirstprivateOperandsMutable().append(op.getResult());
2522}
2523
2524void acc::ParallelOp::addReduction(MLIRContext *context,
2525 mlir::acc::ReductionOp op,
2526 mlir::acc::ReductionRecipeOp recipe) {
2527 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
2528 getReductionOperandsMutable().append(op.getResult());
2529}
2530
2531static ParseResult parseNumGangs(
2532 mlir::OpAsmParser &parser,
2534 llvm::SmallVectorImpl<Type> &types, mlir::ArrayAttr &deviceTypes,
2535 mlir::DenseI32ArrayAttr &segments) {
2538
2539 do {
2540 if (failed(parser.parseLBrace()))
2541 return failure();
2542
2543 int32_t crtOperandsSize = operands.size();
2544 if (failed(parser.parseCommaSeparatedList(
2546 if (parser.parseOperand(operands.emplace_back()) ||
2547 parser.parseColonType(types.emplace_back()))
2548 return failure();
2549 return success();
2550 })))
2551 return failure();
2552 seg.push_back(operands.size() - crtOperandsSize);
2553
2554 if (failed(parser.parseRBrace()))
2555 return failure();
2556
2557 if (succeeded(parser.parseOptionalLSquare())) {
2558 if (parser.parseAttribute(attributes.emplace_back()) ||
2559 parser.parseRSquare())
2560 return failure();
2561 } else {
2562 attributes.push_back(mlir::acc::DeviceTypeAttr::get(
2563 parser.getContext(), mlir::acc::DeviceType::None));
2564 }
2565 } while (succeeded(parser.parseOptionalComma()));
2566
2567 llvm::SmallVector<mlir::Attribute> arrayAttr(attributes.begin(),
2568 attributes.end());
2569 deviceTypes = ArrayAttr::get(parser.getContext(), arrayAttr);
2570 segments = DenseI32ArrayAttr::get(parser.getContext(), seg);
2571
2572 return success();
2573}
2574
2576 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
2577 if (deviceTypeAttr.getValue() != mlir::acc::DeviceType::None)
2578 p << " [" << attr << "]";
2579}
2580
2582 mlir::OperandRange operands, mlir::TypeRange types,
2583 std::optional<mlir::ArrayAttr> deviceTypes,
2584 std::optional<mlir::DenseI32ArrayAttr> segments) {
2585 unsigned opIdx = 0;
2586 llvm::interleaveComma(llvm::enumerate(*deviceTypes), p, [&](auto it) {
2587 p << "{";
2588 llvm::interleaveComma(
2589 llvm::seq<int32_t>(0, (*segments)[it.index()]), p, [&](auto it) {
2590 p << operands[opIdx] << " : " << operands[opIdx].getType();
2591 ++opIdx;
2592 });
2593 p << "}";
2594 printSingleDeviceType(p, it.value());
2595 });
2596}
2597
2599 mlir::OpAsmParser &parser,
2601 llvm::SmallVectorImpl<Type> &types, mlir::ArrayAttr &deviceTypes,
2602 mlir::DenseI32ArrayAttr &segments) {
2605
2606 do {
2607 if (failed(parser.parseLBrace()))
2608 return failure();
2609
2610 int32_t crtOperandsSize = operands.size();
2611
2612 if (failed(parser.parseCommaSeparatedList(
2614 if (parser.parseOperand(operands.emplace_back()) ||
2615 parser.parseColonType(types.emplace_back()))
2616 return failure();
2617 return success();
2618 })))
2619 return failure();
2620
2621 seg.push_back(operands.size() - crtOperandsSize);
2622
2623 if (failed(parser.parseRBrace()))
2624 return failure();
2625
2626 if (succeeded(parser.parseOptionalLSquare())) {
2627 if (parser.parseAttribute(attributes.emplace_back()) ||
2628 parser.parseRSquare())
2629 return failure();
2630 } else {
2631 attributes.push_back(mlir::acc::DeviceTypeAttr::get(
2632 parser.getContext(), mlir::acc::DeviceType::None));
2633 }
2634 } while (succeeded(parser.parseOptionalComma()));
2635
2636 llvm::SmallVector<mlir::Attribute> arrayAttr(attributes.begin(),
2637 attributes.end());
2638 deviceTypes = ArrayAttr::get(parser.getContext(), arrayAttr);
2639 segments = DenseI32ArrayAttr::get(parser.getContext(), seg);
2640
2641 return success();
2642}
2643
2646 mlir::TypeRange types, std::optional<mlir::ArrayAttr> deviceTypes,
2647 std::optional<mlir::DenseI32ArrayAttr> segments) {
2648 unsigned opIdx = 0;
2649 llvm::interleaveComma(llvm::enumerate(*deviceTypes), p, [&](auto it) {
2650 p << "{";
2651 llvm::interleaveComma(
2652 llvm::seq<int32_t>(0, (*segments)[it.index()]), p, [&](auto it) {
2653 p << operands[opIdx] << " : " << operands[opIdx].getType();
2654 ++opIdx;
2655 });
2656 p << "}";
2657 printSingleDeviceType(p, it.value());
2658 });
2659}
2660
2661static ParseResult parseWaitClause(
2662 mlir::OpAsmParser &parser,
2664 llvm::SmallVectorImpl<Type> &types, mlir::ArrayAttr &deviceTypes,
2665 mlir::DenseI32ArrayAttr &segments, mlir::ArrayAttr &hasDevNum,
2666 mlir::ArrayAttr &keywordOnly) {
2667 llvm::SmallVector<mlir::Attribute> deviceTypeAttrs, keywordAttrs, devnum;
2669
2670 bool needCommaBeforeOperands = false;
2671
2672 // Keyword only
2673 if (failed(parser.parseOptionalLParen())) {
2674 keywordAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
2675 parser.getContext(), mlir::acc::DeviceType::None));
2676 keywordOnly = ArrayAttr::get(parser.getContext(), keywordAttrs);
2677 return success();
2678 }
2679
2680 // Parse keyword only attributes
2681 if (succeeded(parser.parseOptionalLSquare())) {
2682 if (failed(parser.parseCommaSeparatedList([&]() {
2683 if (parser.parseAttribute(keywordAttrs.emplace_back()))
2684 return failure();
2685 return success();
2686 })))
2687 return failure();
2688 if (parser.parseRSquare())
2689 return failure();
2690 needCommaBeforeOperands = true;
2691 }
2692
2693 if (needCommaBeforeOperands && failed(parser.parseComma()))
2694 return failure();
2695
2696 do {
2697 if (failed(parser.parseLBrace()))
2698 return failure();
2699
2700 int32_t crtOperandsSize = operands.size();
2701
2702 if (succeeded(parser.parseOptionalKeyword("devnum"))) {
2703 if (failed(parser.parseColon()))
2704 return failure();
2705 devnum.push_back(BoolAttr::get(parser.getContext(), true));
2706 } else {
2707 devnum.push_back(BoolAttr::get(parser.getContext(), false));
2708 }
2709
2710 if (failed(parser.parseCommaSeparatedList(
2712 if (parser.parseOperand(operands.emplace_back()) ||
2713 parser.parseColonType(types.emplace_back()))
2714 return failure();
2715 return success();
2716 })))
2717 return failure();
2718
2719 seg.push_back(operands.size() - crtOperandsSize);
2720
2721 if (failed(parser.parseRBrace()))
2722 return failure();
2723
2724 if (succeeded(parser.parseOptionalLSquare())) {
2725 if (parser.parseAttribute(deviceTypeAttrs.emplace_back()) ||
2726 parser.parseRSquare())
2727 return failure();
2728 } else {
2729 deviceTypeAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
2730 parser.getContext(), mlir::acc::DeviceType::None));
2731 }
2732 } while (succeeded(parser.parseOptionalComma()));
2733
2734 if (failed(parser.parseRParen()))
2735 return failure();
2736
2737 deviceTypes = ArrayAttr::get(parser.getContext(), deviceTypeAttrs);
2738 keywordOnly = ArrayAttr::get(parser.getContext(), keywordAttrs);
2739 segments = DenseI32ArrayAttr::get(parser.getContext(), seg);
2740 hasDevNum = ArrayAttr::get(parser.getContext(), devnum);
2741
2742 return success();
2743}
2744
2745static bool hasOnlyDeviceTypeNone(std::optional<mlir::ArrayAttr> attrs) {
2746 if (!hasDeviceTypeValues(attrs))
2747 return false;
2748 if (attrs->size() != 1)
2749 return false;
2750 if (auto deviceTypeAttr =
2751 mlir::dyn_cast<mlir::acc::DeviceTypeAttr>((*attrs)[0]))
2752 return deviceTypeAttr.getValue() == mlir::acc::DeviceType::None;
2753 return false;
2754}
2755
2757 mlir::OperandRange operands, mlir::TypeRange types,
2758 std::optional<mlir::ArrayAttr> deviceTypes,
2759 std::optional<mlir::DenseI32ArrayAttr> segments,
2760 std::optional<mlir::ArrayAttr> hasDevNum,
2761 std::optional<mlir::ArrayAttr> keywordOnly) {
2762
2763 if (operands.begin() == operands.end() && hasOnlyDeviceTypeNone(keywordOnly))
2764 return;
2765
2766 p << "(";
2767
2768 printDeviceTypes(p, keywordOnly);
2769 if (hasDeviceTypeValues(keywordOnly) && hasDeviceTypeValues(deviceTypes))
2770 p << ", ";
2771
2772 if (hasDeviceTypeValues(deviceTypes)) {
2773 unsigned opIdx = 0;
2774 llvm::interleaveComma(llvm::enumerate(*deviceTypes), p, [&](auto it) {
2775 p << "{";
2776 auto boolAttr = mlir::dyn_cast<mlir::BoolAttr>((*hasDevNum)[it.index()]);
2777 if (boolAttr && boolAttr.getValue())
2778 p << "devnum: ";
2779 llvm::interleaveComma(
2780 llvm::seq<int32_t>(0, (*segments)[it.index()]), p, [&](auto it) {
2781 p << operands[opIdx] << " : " << operands[opIdx].getType();
2782 ++opIdx;
2783 });
2784 p << "}";
2785 printSingleDeviceType(p, it.value());
2786 });
2787 }
2788
2789 p << ")";
2790}
2791
2792static ParseResult parseDeviceTypeOperands(
2793 mlir::OpAsmParser &parser,
2795 llvm::SmallVectorImpl<Type> &types, mlir::ArrayAttr &deviceTypes) {
2797 if (failed(parser.parseCommaSeparatedList([&]() {
2798 if (parser.parseOperand(operands.emplace_back()) ||
2799 parser.parseColonType(types.emplace_back()))
2800 return failure();
2801 if (succeeded(parser.parseOptionalLSquare())) {
2802 if (parser.parseAttribute(attributes.emplace_back()) ||
2803 parser.parseRSquare())
2804 return failure();
2805 } else {
2806 attributes.push_back(mlir::acc::DeviceTypeAttr::get(
2807 parser.getContext(), mlir::acc::DeviceType::None));
2808 }
2809 return success();
2810 })))
2811 return failure();
2812 llvm::SmallVector<mlir::Attribute> arrayAttr(attributes.begin(),
2813 attributes.end());
2814 deviceTypes = ArrayAttr::get(parser.getContext(), arrayAttr);
2815 return success();
2816}
2817
2818static void
2820 mlir::OperandRange operands, mlir::TypeRange types,
2821 std::optional<mlir::ArrayAttr> deviceTypes) {
2822 if (!hasDeviceTypeValues(deviceTypes))
2823 return;
2824 llvm::interleaveComma(llvm::zip(*deviceTypes, operands), p, [&](auto it) {
2825 p << std::get<1>(it) << " : " << std::get<1>(it).getType();
2826 printSingleDeviceType(p, std::get<0>(it));
2827 });
2828}
2829
2831 mlir::OpAsmParser &parser,
2833 llvm::SmallVectorImpl<Type> &types, mlir::ArrayAttr &deviceTypes,
2834 mlir::ArrayAttr &keywordOnlyDeviceType) {
2835
2836 llvm::SmallVector<mlir::Attribute> keywordOnlyDeviceTypeAttributes;
2837 bool needCommaBeforeOperands = false;
2838
2839 if (failed(parser.parseOptionalLParen())) {
2840 // Keyword only
2841 keywordOnlyDeviceTypeAttributes.push_back(mlir::acc::DeviceTypeAttr::get(
2842 parser.getContext(), mlir::acc::DeviceType::None));
2843 keywordOnlyDeviceType =
2844 ArrayAttr::get(parser.getContext(), keywordOnlyDeviceTypeAttributes);
2845 return success();
2846 }
2847
2848 // Parse keyword only attributes
2849 if (succeeded(parser.parseOptionalLSquare())) {
2850 // Parse keyword only attributes
2851 if (failed(parser.parseCommaSeparatedList([&]() {
2852 if (parser.parseAttribute(
2853 keywordOnlyDeviceTypeAttributes.emplace_back()))
2854 return failure();
2855 return success();
2856 })))
2857 return failure();
2858 if (parser.parseRSquare())
2859 return failure();
2860 keywordOnlyDeviceType =
2861 ArrayAttr::get(parser.getContext(), keywordOnlyDeviceTypeAttributes);
2862 needCommaBeforeOperands = true;
2863 }
2864
2865 if (needCommaBeforeOperands) {
2866 if (succeeded(parser.parseOptionalRParen()))
2867 return success();
2868 if (failed(parser.parseComma()))
2869 return failure();
2870 }
2871
2873 if (failed(parser.parseCommaSeparatedList([&]() {
2874 if (parser.parseOperand(operands.emplace_back()) ||
2875 parser.parseColonType(types.emplace_back()))
2876 return failure();
2877 if (succeeded(parser.parseOptionalLSquare())) {
2878 if (parser.parseAttribute(attributes.emplace_back()) ||
2879 parser.parseRSquare())
2880 return failure();
2881 } else {
2882 attributes.push_back(mlir::acc::DeviceTypeAttr::get(
2883 parser.getContext(), mlir::acc::DeviceType::None));
2884 }
2885 return success();
2886 })))
2887 return failure();
2888
2889 if (failed(parser.parseRParen()))
2890 return failure();
2891
2892 llvm::SmallVector<mlir::Attribute> arrayAttr(attributes.begin(),
2893 attributes.end());
2894 deviceTypes = ArrayAttr::get(parser.getContext(), arrayAttr);
2895 return success();
2896}
2897
2900 mlir::TypeRange types, std::optional<mlir::ArrayAttr> deviceTypes,
2901 std::optional<mlir::ArrayAttr> keywordOnlyDeviceTypes) {
2902
2903 if (operands.begin() == operands.end() &&
2904 hasOnlyDeviceTypeNone(keywordOnlyDeviceTypes)) {
2905 return;
2906 }
2907
2908 p << "(";
2909 printDeviceTypes(p, keywordOnlyDeviceTypes);
2910 if (hasDeviceTypeValues(keywordOnlyDeviceTypes) &&
2911 hasDeviceTypeValues(deviceTypes))
2912 p << ", ";
2913 printDeviceTypeOperands(p, op, operands, types, deviceTypes);
2914 p << ")";
2915}
2916
2918 mlir::OpAsmParser &parser,
2919 std::optional<OpAsmParser::UnresolvedOperand> &operand,
2920 mlir::Type &operandType, mlir::UnitAttr &attr) {
2921 // Keyword only
2922 if (failed(parser.parseOptionalLParen())) {
2923 attr = mlir::UnitAttr::get(parser.getContext());
2924 return success();
2925 }
2926
2928 if (failed(parser.parseOperand(op)))
2929 return failure();
2930 operand = op;
2931 if (failed(parser.parseColon()))
2932 return failure();
2933 if (failed(parser.parseType(operandType)))
2934 return failure();
2935 if (failed(parser.parseRParen()))
2936 return failure();
2937
2938 return success();
2939}
2940
2942 mlir::Operation *op,
2943 std::optional<mlir::Value> operand,
2944 mlir::Type operandType,
2945 mlir::UnitAttr attr) {
2946 if (attr)
2947 return;
2948
2949 p << "(";
2950 p.printOperand(*operand);
2951 p << " : ";
2952 p.printType(operandType);
2953 p << ")";
2954}
2955
2957 mlir::OpAsmParser &parser,
2959 llvm::SmallVectorImpl<Type> &types, mlir::UnitAttr &attr) {
2960 // Keyword only
2961 if (failed(parser.parseOptionalLParen())) {
2962 attr = mlir::UnitAttr::get(parser.getContext());
2963 return success();
2964 }
2965
2966 if (failed(parser.parseCommaSeparatedList([&]() {
2967 if (parser.parseOperand(operands.emplace_back()))
2968 return failure();
2969 return success();
2970 })))
2971 return failure();
2972 if (failed(parser.parseColon()))
2973 return failure();
2974 if (failed(parser.parseCommaSeparatedList([&]() {
2975 if (parser.parseType(types.emplace_back()))
2976 return failure();
2977 return success();
2978 })))
2979 return failure();
2980 if (failed(parser.parseRParen()))
2981 return failure();
2982
2983 return success();
2984}
2985
2987 mlir::Operation *op,
2988 mlir::OperandRange operands,
2989 mlir::TypeRange types,
2990 mlir::UnitAttr attr) {
2991 if (attr)
2992 return;
2993
2994 p << "(";
2995 llvm::interleaveComma(operands, p, [&](auto it) { p << it; });
2996 p << " : ";
2997 llvm::interleaveComma(types, p, [&](auto it) { p << it; });
2998 p << ")";
2999}
3000
3001static ParseResult
3003 mlir::acc::CombinedConstructsTypeAttr &attr) {
3004 if (succeeded(parser.parseOptionalKeyword("kernels"))) {
3005 attr = mlir::acc::CombinedConstructsTypeAttr::get(
3006 parser.getContext(), mlir::acc::CombinedConstructsType::KernelsLoop);
3007 } else if (succeeded(parser.parseOptionalKeyword("parallel"))) {
3008 attr = mlir::acc::CombinedConstructsTypeAttr::get(
3009 parser.getContext(), mlir::acc::CombinedConstructsType::ParallelLoop);
3010 } else if (succeeded(parser.parseOptionalKeyword("serial"))) {
3011 attr = mlir::acc::CombinedConstructsTypeAttr::get(
3012 parser.getContext(), mlir::acc::CombinedConstructsType::SerialLoop);
3013 } else {
3014 parser.emitError(parser.getCurrentLocation(),
3015 "expected compute construct name");
3016 return failure();
3017 }
3018 return success();
3019}
3020
3021static void
3023 mlir::acc::CombinedConstructsTypeAttr attr) {
3024 if (attr) {
3025 switch (attr.getValue()) {
3026 case mlir::acc::CombinedConstructsType::KernelsLoop:
3027 p << "kernels";
3028 break;
3029 case mlir::acc::CombinedConstructsType::ParallelLoop:
3030 p << "parallel";
3031 break;
3032 case mlir::acc::CombinedConstructsType::SerialLoop:
3033 p << "serial";
3034 break;
3035 };
3036 }
3037}
3038
3039//===----------------------------------------------------------------------===//
3040// SerialOp
3041//===----------------------------------------------------------------------===//
3042
3043unsigned SerialOp::getNumDataOperands() {
3044 return getReductionOperands().size() + getPrivateOperands().size() +
3045 getFirstprivateOperands().size() + getDataClauseOperands().size();
3046}
3047
3048Value SerialOp::getDataOperand(unsigned i) {
3049 unsigned numOptional = getAsyncOperands().size();
3050 numOptional += getIfCond() ? 1 : 0;
3051 numOptional += getSelfCond() ? 1 : 0;
3052 return getOperand(getWaitOperands().size() + numOptional + i);
3053}
3054
3055bool acc::SerialOp::hasAsyncOnly() {
3056 return hasAsyncOnly(mlir::acc::DeviceType::None);
3057}
3058
3059bool acc::SerialOp::hasAsyncOnly(mlir::acc::DeviceType deviceType) {
3060 return hasDeviceType(getAsyncOnly(), deviceType);
3061}
3062
3063mlir::Value acc::SerialOp::getAsyncValue() {
3064 return getAsyncValue(mlir::acc::DeviceType::None);
3065}
3066
3067mlir::Value acc::SerialOp::getAsyncValue(mlir::acc::DeviceType deviceType) {
3069 getAsyncOperands(), deviceType);
3070}
3071
3072bool acc::SerialOp::hasWaitOnly() {
3073 return hasWaitOnly(mlir::acc::DeviceType::None);
3074}
3075
3076bool acc::SerialOp::hasWaitOnly(mlir::acc::DeviceType deviceType) {
3077 return hasDeviceType(getWaitOnly(), deviceType);
3078}
3079
3080mlir::Operation::operand_range SerialOp::getWaitValues() {
3081 return getWaitValues(mlir::acc::DeviceType::None);
3082}
3083
3085SerialOp::getWaitValues(mlir::acc::DeviceType deviceType) {
3087 getWaitOperandsDeviceType(), getWaitOperands(), getWaitOperandsSegments(),
3088 getHasWaitDevnum(), deviceType);
3089}
3090
3091mlir::Value SerialOp::getWaitDevnum() {
3092 return getWaitDevnum(mlir::acc::DeviceType::None);
3093}
3094
3095mlir::Value SerialOp::getWaitDevnum(mlir::acc::DeviceType deviceType) {
3096 return getWaitDevnumValue(getWaitOperandsDeviceType(), getWaitOperands(),
3097 getWaitOperandsSegments(), getHasWaitDevnum(),
3098 deviceType);
3099}
3100
3101LogicalResult acc::SerialOp::verify() {
3102 if (failed(checkPrivateOperands<mlir::acc::PrivateOp,
3103 mlir::acc::PrivateRecipeOp>(
3104 *this, getPrivateOperands(), "private")))
3105 return failure();
3106 if (failed(checkPrivateOperands<mlir::acc::FirstprivateOp,
3107 mlir::acc::FirstprivateRecipeOp>(
3108 *this, getFirstprivateOperands(), "firstprivate")))
3109 return failure();
3110 if (failed(checkPrivateOperands<mlir::acc::ReductionOp,
3111 mlir::acc::ReductionRecipeOp>(
3112 *this, getReductionOperands(), "reduction")))
3113 return failure();
3114
3116 *this, getWaitOperands(), getWaitOperandsSegmentsAttr(),
3117 getWaitOperandsDeviceTypeAttr(), "wait")))
3118 return failure();
3119
3121 getAsyncOperandsDeviceTypeAttr(),
3122 "async")))
3123 return failure();
3124
3126 return failure();
3127
3128 return checkDataOperands<acc::SerialOp>(*this, getDataClauseOperands());
3129}
3130
3131void acc::SerialOp::addAsyncOnly(
3132 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3133 setAsyncOnlyAttr(addDeviceTypeAffectedOperandHelper(
3134 context, getAsyncOnlyAttr(), effectiveDeviceTypes));
3135}
3136
3137void acc::SerialOp::addAsyncOperand(
3138 MLIRContext *context, mlir::Value newValue,
3139 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3140 setAsyncOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3141 context, getAsyncOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValue,
3142 getAsyncOperandsMutable()));
3143}
3144
3145void acc::SerialOp::addWaitOnly(
3146 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3147 setWaitOnlyAttr(addDeviceTypeAffectedOperandHelper(context, getWaitOnlyAttr(),
3148 effectiveDeviceTypes));
3149}
3150void acc::SerialOp::addWaitOperands(
3151 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
3152 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3153
3155 if (getWaitOperandsSegments())
3156 llvm::copy(*getWaitOperandsSegments(), std::back_inserter(segments));
3157
3158 setWaitOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3159 context, getWaitOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
3160 getWaitOperandsMutable(), segments));
3161 setWaitOperandsSegments(segments);
3162
3164 if (getHasWaitDevnumAttr())
3165 llvm::copy(getHasWaitDevnumAttr(), std::back_inserter(hasDevnums));
3166 hasDevnums.insert(
3167 hasDevnums.end(),
3168 std::max(effectiveDeviceTypes.size(), static_cast<size_t>(1)),
3169 mlir::BoolAttr::get(context, hasDevnum));
3170 setHasWaitDevnumAttr(mlir::ArrayAttr::get(context, hasDevnums));
3171}
3172
3173void acc::SerialOp::addPrivatization(MLIRContext *context,
3174 mlir::acc::PrivateOp op,
3175 mlir::acc::PrivateRecipeOp recipe) {
3176 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3177 getPrivateOperandsMutable().append(op.getResult());
3178}
3179
3180void acc::SerialOp::addFirstPrivatization(
3181 MLIRContext *context, mlir::acc::FirstprivateOp op,
3182 mlir::acc::FirstprivateRecipeOp recipe) {
3183 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3184 getFirstprivateOperandsMutable().append(op.getResult());
3185}
3186
3187void acc::SerialOp::addReduction(MLIRContext *context,
3188 mlir::acc::ReductionOp op,
3189 mlir::acc::ReductionRecipeOp recipe) {
3190 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3191 getReductionOperandsMutable().append(op.getResult());
3192}
3193
3194//===----------------------------------------------------------------------===//
3195// KernelsOp
3196//===----------------------------------------------------------------------===//
3197
3198unsigned KernelsOp::getNumDataOperands() {
3199 return getDataClauseOperands().size();
3200}
3201
3202Value KernelsOp::getDataOperand(unsigned i) {
3203 unsigned numOptional = getAsyncOperands().size();
3204 numOptional += getWaitOperands().size();
3205 numOptional += getNumGangs().size();
3206 numOptional += getNumWorkers().size();
3207 numOptional += getVectorLength().size();
3208 numOptional += getIfCond() ? 1 : 0;
3209 numOptional += getSelfCond() ? 1 : 0;
3210 return getOperand(numOptional + i);
3211}
3212
3213bool acc::KernelsOp::hasAsyncOnly() {
3214 return hasAsyncOnly(mlir::acc::DeviceType::None);
3215}
3216
3217bool acc::KernelsOp::hasAsyncOnly(mlir::acc::DeviceType deviceType) {
3218 return hasDeviceType(getAsyncOnly(), deviceType);
3219}
3220
3221mlir::Value acc::KernelsOp::getAsyncValue() {
3222 return getAsyncValue(mlir::acc::DeviceType::None);
3223}
3224
3225mlir::Value acc::KernelsOp::getAsyncValue(mlir::acc::DeviceType deviceType) {
3227 getAsyncOperands(), deviceType);
3228}
3229
3230mlir::Value acc::KernelsOp::getNumWorkersValue() {
3231 return getNumWorkersValue(mlir::acc::DeviceType::None);
3232}
3233
3235acc::KernelsOp::getNumWorkersValue(mlir::acc::DeviceType deviceType) {
3236 return getValueInDeviceTypeSegment(getNumWorkersDeviceType(), getNumWorkers(),
3237 deviceType);
3238}
3239
3240mlir::Value acc::KernelsOp::getVectorLengthValue() {
3241 return getVectorLengthValue(mlir::acc::DeviceType::None);
3242}
3243
3245acc::KernelsOp::getVectorLengthValue(mlir::acc::DeviceType deviceType) {
3246 return getValueInDeviceTypeSegment(getVectorLengthDeviceType(),
3247 getVectorLength(), deviceType);
3248}
3249
3250mlir::Operation::operand_range KernelsOp::getNumGangsValues() {
3251 return getNumGangsValues(mlir::acc::DeviceType::None);
3252}
3253
3255KernelsOp::getNumGangsValues(mlir::acc::DeviceType deviceType) {
3256 return getValuesFromSegments(getNumGangsDeviceType(), getNumGangs(),
3257 getNumGangsSegments(), deviceType);
3258}
3259
3260bool acc::KernelsOp::hasAnyGangWorkerVector(mlir::acc::DeviceType deviceType) {
3262 getNumGangsDeviceType(), getNumGangs(), getNumGangsSegments(),
3263 getNumWorkersDeviceType(), getNumWorkers(), getVectorLengthDeviceType(),
3264 getVectorLength(), deviceType);
3265}
3266
3267bool acc::KernelsOp::isEffectivelySerial() {
3268 return isGangWorkerVectorAllOne(*this);
3269}
3270
3271bool acc::KernelsOp::hasWaitOnly() {
3272 return hasWaitOnly(mlir::acc::DeviceType::None);
3273}
3274
3275bool acc::KernelsOp::hasWaitOnly(mlir::acc::DeviceType deviceType) {
3276 return hasDeviceType(getWaitOnly(), deviceType);
3277}
3278
3279mlir::Operation::operand_range KernelsOp::getWaitValues() {
3280 return getWaitValues(mlir::acc::DeviceType::None);
3281}
3282
3284KernelsOp::getWaitValues(mlir::acc::DeviceType deviceType) {
3286 getWaitOperandsDeviceType(), getWaitOperands(), getWaitOperandsSegments(),
3287 getHasWaitDevnum(), deviceType);
3288}
3289
3290mlir::Value KernelsOp::getWaitDevnum() {
3291 return getWaitDevnum(mlir::acc::DeviceType::None);
3292}
3293
3294mlir::Value KernelsOp::getWaitDevnum(mlir::acc::DeviceType deviceType) {
3295 return getWaitDevnumValue(getWaitOperandsDeviceType(), getWaitOperands(),
3296 getWaitOperandsSegments(), getHasWaitDevnum(),
3297 deviceType);
3298}
3299
3300LogicalResult acc::KernelsOp::verify() {
3302 *this, getNumGangs(), getNumGangsSegmentsAttr(),
3303 getNumGangsDeviceTypeAttr(), "num_gangs", 3)))
3304 return failure();
3305
3307 *this, getWaitOperands(), getWaitOperandsSegmentsAttr(),
3308 getWaitOperandsDeviceTypeAttr(), "wait")))
3309 return failure();
3310
3311 if (failed(verifyDeviceTypeCountMatch(*this, getNumWorkers(),
3312 getNumWorkersDeviceTypeAttr(),
3313 "num_workers")))
3314 return failure();
3315
3316 if (failed(verifyDeviceTypeCountMatch(*this, getVectorLength(),
3317 getVectorLengthDeviceTypeAttr(),
3318 "vector_length")))
3319 return failure();
3320
3322 getAsyncOperandsDeviceTypeAttr(),
3323 "async")))
3324 return failure();
3325
3327 return failure();
3328
3329 return checkDataOperands<acc::KernelsOp>(*this, getDataClauseOperands());
3330}
3331
3332void acc::KernelsOp::addPrivatization(MLIRContext *context,
3333 mlir::acc::PrivateOp op,
3334 mlir::acc::PrivateRecipeOp recipe) {
3335 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3336 getPrivateOperandsMutable().append(op.getResult());
3337}
3338
3339void acc::KernelsOp::addFirstPrivatization(
3340 MLIRContext *context, mlir::acc::FirstprivateOp op,
3341 mlir::acc::FirstprivateRecipeOp recipe) {
3342 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3343 getFirstprivateOperandsMutable().append(op.getResult());
3344}
3345
3346void acc::KernelsOp::addReduction(MLIRContext *context,
3347 mlir::acc::ReductionOp op,
3348 mlir::acc::ReductionRecipeOp recipe) {
3349 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
3350 getReductionOperandsMutable().append(op.getResult());
3351}
3352
3353void acc::KernelsOp::addNumWorkersOperand(
3354 MLIRContext *context, mlir::Value newValue,
3355 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3356 setNumWorkersDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3357 context, getNumWorkersDeviceTypeAttr(), effectiveDeviceTypes, newValue,
3358 getNumWorkersMutable()));
3359}
3360
3361void acc::KernelsOp::addVectorLengthOperand(
3362 MLIRContext *context, mlir::Value newValue,
3363 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3364 setVectorLengthDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3365 context, getVectorLengthDeviceTypeAttr(), effectiveDeviceTypes, newValue,
3366 getVectorLengthMutable()));
3367}
3368void acc::KernelsOp::addAsyncOnly(
3369 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3370 setAsyncOnlyAttr(addDeviceTypeAffectedOperandHelper(
3371 context, getAsyncOnlyAttr(), effectiveDeviceTypes));
3372}
3373
3374void acc::KernelsOp::addAsyncOperand(
3375 MLIRContext *context, mlir::Value newValue,
3376 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3377 setAsyncOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3378 context, getAsyncOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValue,
3379 getAsyncOperandsMutable()));
3380}
3381
3382void acc::KernelsOp::addNumGangsOperands(
3383 MLIRContext *context, mlir::ValueRange newValues,
3384 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3386 if (getNumGangsSegmentsAttr())
3387 llvm::copy(*getNumGangsSegments(), std::back_inserter(segments));
3388
3389 setNumGangsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3390 context, getNumGangsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
3391 getNumGangsMutable(), segments));
3392
3393 setNumGangsSegments(segments);
3394}
3395
3396void acc::KernelsOp::addWaitOnly(
3397 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3398 setWaitOnlyAttr(addDeviceTypeAffectedOperandHelper(context, getWaitOnlyAttr(),
3399 effectiveDeviceTypes));
3400}
3401void acc::KernelsOp::addWaitOperands(
3402 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
3403 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
3404
3406 if (getWaitOperandsSegments())
3407 llvm::copy(*getWaitOperandsSegments(), std::back_inserter(segments));
3408
3409 setWaitOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
3410 context, getWaitOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
3411 getWaitOperandsMutable(), segments));
3412 setWaitOperandsSegments(segments);
3413
3415 if (getHasWaitDevnumAttr())
3416 llvm::copy(getHasWaitDevnumAttr(), std::back_inserter(hasDevnums));
3417 hasDevnums.insert(
3418 hasDevnums.end(),
3419 std::max(effectiveDeviceTypes.size(), static_cast<size_t>(1)),
3420 mlir::BoolAttr::get(context, hasDevnum));
3421 setHasWaitDevnumAttr(mlir::ArrayAttr::get(context, hasDevnums));
3422}
3423
3424//===----------------------------------------------------------------------===//
3425// HostDataOp
3426//===----------------------------------------------------------------------===//
3427
3428LogicalResult acc::HostDataOp::verify() {
3429 if (getDataClauseOperands().empty())
3430 return emitError("at least one operand must appear on the host_data "
3431 "operation");
3432
3434 for (mlir::Value operand : getDataClauseOperands()) {
3435 auto useDeviceOp =
3436 mlir::dyn_cast_if_present<acc::UseDeviceOp>(operand.getDefiningOp());
3437 if (!useDeviceOp)
3438 return emitError("expect data entry operation as defining op");
3439
3440 // Check for duplicate use_device clauses
3441 if (!seenVars.insert(useDeviceOp.getVar()).second)
3442 return emitError("duplicate use_device variable");
3443 }
3444 return success();
3445}
3446
3447void acc::HostDataOp::getCanonicalizationPatterns(RewritePatternSet &results,
3448 MLIRContext *context) {
3449 results.add<RemoveConstantIfConditionWithRegion<HostDataOp>>(context);
3450}
3451
3452//===----------------------------------------------------------------------===//
3453// LoopOp
3454//===----------------------------------------------------------------------===//
3455
3456static ParseResult parseGangValue(
3457 OpAsmParser &parser, llvm::StringRef keyword,
3460 llvm::SmallVector<GangArgTypeAttr> &attributes, GangArgTypeAttr gangArgType,
3461 bool &needCommaBetweenValues, bool &newValue) {
3462 if (succeeded(parser.parseOptionalKeyword(keyword))) {
3463 if (parser.parseEqual())
3464 return failure();
3465 if (parser.parseOperand(operands.emplace_back()) ||
3466 parser.parseColonType(types.emplace_back()))
3467 return failure();
3468 attributes.push_back(gangArgType);
3469 needCommaBetweenValues = true;
3470 newValue = true;
3471 }
3472 return success();
3473}
3474
3475static ParseResult parseGangClause(
3476 OpAsmParser &parser,
3478 llvm::SmallVectorImpl<Type> &gangOperandsType, mlir::ArrayAttr &gangArgType,
3479 mlir::ArrayAttr &deviceType, mlir::DenseI32ArrayAttr &segments,
3480 mlir::ArrayAttr &gangOnlyDeviceType) {
3481 llvm::SmallVector<GangArgTypeAttr> gangArgTypeAttributes;
3482 llvm::SmallVector<mlir::Attribute> deviceTypeAttributes;
3483 llvm::SmallVector<mlir::Attribute> gangOnlyDeviceTypeAttributes;
3485 bool needCommaBetweenValues = false;
3486 bool needCommaBeforeOperands = false;
3487
3488 if (failed(parser.parseOptionalLParen())) {
3489 // Gang only keyword
3490 gangOnlyDeviceTypeAttributes.push_back(mlir::acc::DeviceTypeAttr::get(
3491 parser.getContext(), mlir::acc::DeviceType::None));
3492 gangOnlyDeviceType =
3493 ArrayAttr::get(parser.getContext(), gangOnlyDeviceTypeAttributes);
3494 return success();
3495 }
3496
3497 // Parse gang only attributes
3498 if (succeeded(parser.parseOptionalLSquare())) {
3499 // Parse gang only attributes
3500 if (failed(parser.parseCommaSeparatedList([&]() {
3501 if (parser.parseAttribute(
3502 gangOnlyDeviceTypeAttributes.emplace_back()))
3503 return failure();
3504 return success();
3505 })))
3506 return failure();
3507 if (parser.parseRSquare())
3508 return failure();
3509 needCommaBeforeOperands = true;
3510 }
3511
3512 auto argNum = mlir::acc::GangArgTypeAttr::get(parser.getContext(),
3513 mlir::acc::GangArgType::Num);
3514 auto argDim = mlir::acc::GangArgTypeAttr::get(parser.getContext(),
3515 mlir::acc::GangArgType::Dim);
3516 auto argStatic = mlir::acc::GangArgTypeAttr::get(
3517 parser.getContext(), mlir::acc::GangArgType::Static);
3518
3519 do {
3520 if (needCommaBeforeOperands) {
3521 needCommaBeforeOperands = false;
3522 continue;
3523 }
3524
3525 if (failed(parser.parseLBrace()))
3526 return failure();
3527
3528 int32_t crtOperandsSize = gangOperands.size();
3529 while (true) {
3530 bool newValue = false;
3531 bool needValue = false;
3532 if (needCommaBetweenValues) {
3533 if (succeeded(parser.parseOptionalComma()))
3534 needValue = true; // expect a new value after comma.
3535 else
3536 break;
3537 }
3538
3539 if (failed(parseGangValue(parser, LoopOp::getGangNumKeyword(),
3540 gangOperands, gangOperandsType,
3541 gangArgTypeAttributes, argNum,
3542 needCommaBetweenValues, newValue)))
3543 return failure();
3544 if (failed(parseGangValue(parser, LoopOp::getGangDimKeyword(),
3545 gangOperands, gangOperandsType,
3546 gangArgTypeAttributes, argDim,
3547 needCommaBetweenValues, newValue)))
3548 return failure();
3549 if (failed(parseGangValue(parser, LoopOp::getGangStaticKeyword(),
3550 gangOperands, gangOperandsType,
3551 gangArgTypeAttributes, argStatic,
3552 needCommaBetweenValues, newValue)))
3553 return failure();
3554
3555 if (!newValue && needValue) {
3556 parser.emitError(parser.getCurrentLocation(),
3557 "new value expected after comma");
3558 return failure();
3559 }
3560
3561 if (!newValue)
3562 break;
3563 }
3564
3565 if (gangOperands.empty())
3566 return parser.emitError(
3567 parser.getCurrentLocation(),
3568 "expect at least one of num, dim or static values");
3569
3570 if (failed(parser.parseRBrace()))
3571 return failure();
3572
3573 if (succeeded(parser.parseOptionalLSquare())) {
3574 if (parser.parseAttribute(deviceTypeAttributes.emplace_back()) ||
3575 parser.parseRSquare())
3576 return failure();
3577 } else {
3578 deviceTypeAttributes.push_back(mlir::acc::DeviceTypeAttr::get(
3579 parser.getContext(), mlir::acc::DeviceType::None));
3580 }
3581
3582 seg.push_back(gangOperands.size() - crtOperandsSize);
3583
3584 } while (succeeded(parser.parseOptionalComma()));
3585
3586 if (failed(parser.parseRParen()))
3587 return failure();
3588
3589 llvm::SmallVector<mlir::Attribute> arrayAttr(gangArgTypeAttributes.begin(),
3590 gangArgTypeAttributes.end());
3591 gangArgType = ArrayAttr::get(parser.getContext(), arrayAttr);
3592 deviceType = ArrayAttr::get(parser.getContext(), deviceTypeAttributes);
3593
3595 gangOnlyDeviceTypeAttributes.begin(), gangOnlyDeviceTypeAttributes.end());
3596 gangOnlyDeviceType = ArrayAttr::get(parser.getContext(), gangOnlyAttr);
3597
3598 segments = DenseI32ArrayAttr::get(parser.getContext(), seg);
3599 return success();
3600}
3601
3603 mlir::OperandRange operands, mlir::TypeRange types,
3604 std::optional<mlir::ArrayAttr> gangArgTypes,
3605 std::optional<mlir::ArrayAttr> deviceTypes,
3606 std::optional<mlir::DenseI32ArrayAttr> segments,
3607 std::optional<mlir::ArrayAttr> gangOnlyDeviceTypes) {
3608
3609 if (operands.begin() == operands.end() &&
3610 hasOnlyDeviceTypeNone(gangOnlyDeviceTypes)) {
3611 return;
3612 }
3613
3614 p << "(";
3615
3616 printDeviceTypes(p, gangOnlyDeviceTypes);
3617
3618 if (hasDeviceTypeValues(gangOnlyDeviceTypes) &&
3619 hasDeviceTypeValues(deviceTypes))
3620 p << ", ";
3621
3622 if (hasDeviceTypeValues(deviceTypes)) {
3623 unsigned opIdx = 0;
3624 llvm::interleaveComma(llvm::enumerate(*deviceTypes), p, [&](auto it) {
3625 p << "{";
3626 llvm::interleaveComma(
3627 llvm::seq<int32_t>(0, (*segments)[it.index()]), p, [&](auto it) {
3628 auto gangArgTypeAttr = mlir::dyn_cast<mlir::acc::GangArgTypeAttr>(
3629 (*gangArgTypes)[opIdx]);
3630 if (gangArgTypeAttr.getValue() == mlir::acc::GangArgType::Num)
3631 p << LoopOp::getGangNumKeyword();
3632 else if (gangArgTypeAttr.getValue() == mlir::acc::GangArgType::Dim)
3633 p << LoopOp::getGangDimKeyword();
3634 else if (gangArgTypeAttr.getValue() ==
3635 mlir::acc::GangArgType::Static)
3636 p << LoopOp::getGangStaticKeyword();
3637 p << "=" << operands[opIdx] << " : " << operands[opIdx].getType();
3638 ++opIdx;
3639 });
3640 p << "}";
3641 printSingleDeviceType(p, it.value());
3642 });
3643 }
3644 p << ")";
3645}
3646
3648 std::optional<mlir::ArrayAttr> segments,
3649 llvm::SmallSet<mlir::acc::DeviceType, 3> &deviceTypes) {
3650 if (!segments)
3651 return false;
3652 for (auto attr : *segments) {
3653 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
3654 if (!deviceTypes.insert(deviceTypeAttr.getValue()).second)
3655 return true;
3656 }
3657 return false;
3658}
3659
3660/// Check for duplicates in the DeviceType array attribute.
3661/// Returns std::nullopt if no duplicates, or the duplicate DeviceType if found.
3662static std::optional<mlir::acc::DeviceType>
3663checkDeviceTypes(mlir::ArrayAttr deviceTypes) {
3664 llvm::SmallSet<mlir::acc::DeviceType, 3> crtDeviceTypes;
3665 if (!deviceTypes)
3666 return std::nullopt;
3667 for (auto attr : deviceTypes) {
3668 auto deviceTypeAttr =
3669 mlir::dyn_cast_or_null<mlir::acc::DeviceTypeAttr>(attr);
3670 if (!deviceTypeAttr)
3671 return mlir::acc::DeviceType::None;
3672 if (!crtDeviceTypes.insert(deviceTypeAttr.getValue()).second)
3673 return deviceTypeAttr.getValue();
3674 }
3675 return std::nullopt;
3676}
3677
3678LogicalResult acc::LoopOp::verify() {
3679 if (getUpperbound().size() != getStep().size())
3680 return emitError() << "number of upperbounds expected to be the same as "
3681 "number of steps";
3682
3683 if (getUpperbound().size() != getLowerbound().size())
3684 return emitError() << "number of upperbounds expected to be the same as "
3685 "number of lowerbounds";
3686
3687 if (!getUpperbound().empty() && getInclusiveUpperbound() &&
3688 (getUpperbound().size() != getInclusiveUpperbound()->size()))
3689 return emitError() << "inclusiveUpperbound size is expected to be the same"
3690 << " as upperbound size";
3691
3692 // Check collapse
3693 if (getCollapseAttr() && !getCollapseDeviceTypeAttr())
3694 return emitOpError() << "collapse device_type attr must be define when"
3695 << " collapse attr is present";
3696
3697 if (getCollapseAttr() && getCollapseDeviceTypeAttr() &&
3698 getCollapseAttr().getValue().size() !=
3699 getCollapseDeviceTypeAttr().getValue().size())
3700 return emitOpError() << "collapse attribute count must match collapse"
3701 << " device_type count";
3702 if (auto duplicateDeviceType = checkDeviceTypes(getCollapseDeviceTypeAttr()))
3703 return emitOpError() << "duplicate device_type `"
3704 << acc::stringifyDeviceType(*duplicateDeviceType)
3705 << "` found in collapseDeviceType attribute";
3706
3707 // Check gang
3708 if (!getGangOperands().empty()) {
3709 if (!getGangOperandsArgType())
3710 return emitOpError() << "gangOperandsArgType attribute must be defined"
3711 << " when gang operands are present";
3712
3713 if (getGangOperands().size() !=
3714 getGangOperandsArgTypeAttr().getValue().size())
3715 return emitOpError() << "gangOperandsArgType attribute count must match"
3716 << " gangOperands count";
3717 }
3718 if (getGangAttr()) {
3719 if (auto duplicateDeviceType = checkDeviceTypes(getGangAttr()))
3720 return emitOpError() << "duplicate device_type `"
3721 << acc::stringifyDeviceType(*duplicateDeviceType)
3722 << "` found in gang attribute";
3723 }
3724
3726 *this, getGangOperands(), getGangOperandsSegmentsAttr(),
3727 getGangOperandsDeviceTypeAttr(), "gang")))
3728 return failure();
3729
3730 // Check worker
3731 if (auto duplicateDeviceType = checkDeviceTypes(getWorkerAttr()))
3732 return emitOpError() << "duplicate device_type `"
3733 << acc::stringifyDeviceType(*duplicateDeviceType)
3734 << "` found in worker attribute";
3735 if (auto duplicateDeviceType =
3736 checkDeviceTypes(getWorkerNumOperandsDeviceTypeAttr()))
3737 return emitOpError() << "duplicate device_type `"
3738 << acc::stringifyDeviceType(*duplicateDeviceType)
3739 << "` found in workerNumOperandsDeviceType attribute";
3740 if (failed(verifyDeviceTypeCountMatch(*this, getWorkerNumOperands(),
3741 getWorkerNumOperandsDeviceTypeAttr(),
3742 "worker")))
3743 return failure();
3744
3745 // Check vector
3746 if (auto duplicateDeviceType = checkDeviceTypes(getVectorAttr()))
3747 return emitOpError() << "duplicate device_type `"
3748 << acc::stringifyDeviceType(*duplicateDeviceType)
3749 << "` found in vector attribute";
3750 if (auto duplicateDeviceType =
3751 checkDeviceTypes(getVectorOperandsDeviceTypeAttr()))
3752 return emitOpError() << "duplicate device_type `"
3753 << acc::stringifyDeviceType(*duplicateDeviceType)
3754 << "` found in vectorOperandsDeviceType attribute";
3755 if (failed(verifyDeviceTypeCountMatch(*this, getVectorOperands(),
3756 getVectorOperandsDeviceTypeAttr(),
3757 "vector")))
3758 return failure();
3759
3761 *this, getTileOperands(), getTileOperandsSegmentsAttr(),
3762 getTileOperandsDeviceTypeAttr(), "tile")))
3763 return failure();
3764
3765 // auto, independent and seq attribute are mutually exclusive.
3766 llvm::SmallSet<mlir::acc::DeviceType, 3> deviceTypes;
3767 if (hasDuplicateDeviceTypes(getAuto_(), deviceTypes) ||
3768 hasDuplicateDeviceTypes(getIndependent(), deviceTypes) ||
3769 hasDuplicateDeviceTypes(getSeq(), deviceTypes)) {
3770 return emitError() << "only one of auto, independent, seq can be present "
3771 "at the same time";
3772 }
3773
3774 // Check that at least one of auto, independent, or seq is present
3775 // for the device-independent default clauses.
3776 auto hasDeviceNone = [](mlir::acc::DeviceTypeAttr attr) -> bool {
3777 return attr.getValue() == mlir::acc::DeviceType::None;
3778 };
3779 bool hasDefaultSeq =
3780 getSeqAttr()
3781 ? llvm::any_of(getSeqAttr().getAsRange<mlir::acc::DeviceTypeAttr>(),
3782 hasDeviceNone)
3783 : false;
3784 bool hasDefaultIndependent =
3785 getIndependentAttr()
3786 ? llvm::any_of(
3787 getIndependentAttr().getAsRange<mlir::acc::DeviceTypeAttr>(),
3788 hasDeviceNone)
3789 : false;
3790 bool hasDefaultAuto =
3791 getAuto_Attr()
3792 ? llvm::any_of(getAuto_Attr().getAsRange<mlir::acc::DeviceTypeAttr>(),
3793 hasDeviceNone)
3794 : false;
3795 if (!hasDefaultSeq && !hasDefaultIndependent && !hasDefaultAuto) {
3796 return emitError()
3797 << "at least one of auto, independent, seq must be present";
3798 }
3799
3800 // Gang, worker and vector are incompatible with seq.
3801 if (getSeqAttr()) {
3802 for (auto attr : getSeqAttr()) {
3803 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
3804 if (hasVector(deviceTypeAttr.getValue()) ||
3805 getVectorValue(deviceTypeAttr.getValue()) ||
3806 hasWorker(deviceTypeAttr.getValue()) ||
3807 getWorkerValue(deviceTypeAttr.getValue()) ||
3808 hasGang(deviceTypeAttr.getValue()) ||
3809 getGangValue(mlir::acc::GangArgType::Num,
3810 deviceTypeAttr.getValue()) ||
3811 getGangValue(mlir::acc::GangArgType::Dim,
3812 deviceTypeAttr.getValue()) ||
3813 getGangValue(mlir::acc::GangArgType::Static,
3814 deviceTypeAttr.getValue()))
3815 return emitError() << "gang, worker or vector cannot appear with seq";
3816 }
3817 }
3818
3819 if (failed(checkPrivateOperands<mlir::acc::PrivateOp,
3820 mlir::acc::PrivateRecipeOp>(
3821 *this, getPrivateOperands(), "private")))
3822 return failure();
3823
3824 if (failed(checkPrivateOperands<mlir::acc::FirstprivateOp,
3825 mlir::acc::FirstprivateRecipeOp>(
3826 *this, getFirstprivateOperands(), "firstprivate")))
3827 return failure();
3828
3829 if (failed(checkPrivateOperands<mlir::acc::ReductionOp,
3830 mlir::acc::ReductionRecipeOp>(
3831 *this, getReductionOperands(), "reduction")))
3832 return failure();
3833
3834 if (getCombined().has_value() &&
3835 (getCombined().value() != acc::CombinedConstructsType::ParallelLoop &&
3836 getCombined().value() != acc::CombinedConstructsType::KernelsLoop &&
3837 getCombined().value() != acc::CombinedConstructsType::SerialLoop)) {
3838 return emitError("unexpected combined constructs attribute");
3839 }
3840
3841 // Check non-empty body().
3842 if (getRegion().empty())
3843 return emitError("expected non-empty body.");
3844
3845 if (getUnstructured()) {
3846 if (!isContainerLike())
3847 return emitError(
3848 "unstructured acc.loop must not have induction variables");
3849 } else if (isContainerLike()) {
3850 // When it is container-like - it is expected to hold a loop-like operation.
3851 // Obtain the maximum collapse count - we use this to check that there
3852 // are enough loops contained.
3853 uint64_t collapseCount = getCollapseValue().value_or(1);
3854 if (getCollapseAttr()) {
3855 for (auto collapseEntry : getCollapseAttr()) {
3856 auto intAttr = mlir::dyn_cast<IntegerAttr>(collapseEntry);
3857 if (intAttr.getValue().getZExtValue() > collapseCount)
3858 collapseCount = intAttr.getValue().getZExtValue();
3859 }
3860 }
3861
3862 // We want to check that we find enough loop-like operations inside.
3863 // PreOrder walk allows us to walk in a breadth-first manner at each nesting
3864 // level.
3865 mlir::Operation *expectedParent = this->getOperation();
3866 bool foundSibling = false;
3867 getRegion().walk<WalkOrder::PreOrder>([&](mlir::Operation *op) {
3868 if (mlir::isa<mlir::LoopLikeOpInterface>(op)) {
3869 // This effectively checks that we are not looking at a sibling loop.
3870 if (op->getParentOfType<mlir::LoopLikeOpInterface>() !=
3871 expectedParent) {
3872 foundSibling = true;
3874 }
3875
3876 collapseCount--;
3877 expectedParent = op;
3878 }
3879 // We found enough contained loops.
3880 if (collapseCount == 0)
3883 });
3884
3885 if (foundSibling)
3886 return emitError("found sibling loops inside container-like acc.loop");
3887 if (collapseCount != 0)
3888 return emitError("failed to find enough loop-like operations inside "
3889 "container-like acc.loop");
3890 }
3891
3892 return success();
3893}
3894
3895unsigned LoopOp::getNumDataOperands() {
3896 return getReductionOperands().size() + getPrivateOperands().size() +
3897 getFirstprivateOperands().size();
3898}
3899
3900Value LoopOp::getDataOperand(unsigned i) {
3901 unsigned numOptional =
3902 getLowerbound().size() + getUpperbound().size() + getStep().size();
3903 numOptional += getGangOperands().size();
3904 numOptional += getVectorOperands().size();
3905 numOptional += getWorkerNumOperands().size();
3906 numOptional += getTileOperands().size();
3907 numOptional += getCacheOperands().size();
3908 return getOperand(numOptional + i);
3909}
3910
3911bool LoopOp::hasAuto() { return hasAuto(mlir::acc::DeviceType::None); }
3912
3913bool LoopOp::hasAuto(mlir::acc::DeviceType deviceType) {
3914 return hasDeviceType(getAuto_(), deviceType);
3915}
3916
3917bool LoopOp::hasIndependent() {
3918 return hasIndependent(mlir::acc::DeviceType::None);
3919}
3920
3921bool LoopOp::hasIndependent(mlir::acc::DeviceType deviceType) {
3922 return hasDeviceType(getIndependent(), deviceType);
3923}
3924
3925bool LoopOp::hasSeq() { return hasSeq(mlir::acc::DeviceType::None); }
3926
3927bool LoopOp::hasSeq(mlir::acc::DeviceType deviceType) {
3928 return hasDeviceType(getSeq(), deviceType);
3929}
3930
3931mlir::Value LoopOp::getVectorValue() {
3932 return getVectorValue(mlir::acc::DeviceType::None);
3933}
3934
3935mlir::Value LoopOp::getVectorValue(mlir::acc::DeviceType deviceType) {
3936 return getValueInDeviceTypeSegment(getVectorOperandsDeviceType(),
3937 getVectorOperands(), deviceType);
3938}
3939
3940bool LoopOp::hasVector() { return hasVector(mlir::acc::DeviceType::None); }
3941
3942bool LoopOp::hasVector(mlir::acc::DeviceType deviceType) {
3943 return hasDeviceType(getVector(), deviceType);
3944}
3945
3946mlir::Value LoopOp::getWorkerValue() {
3947 return getWorkerValue(mlir::acc::DeviceType::None);
3948}
3949
3950mlir::Value LoopOp::getWorkerValue(mlir::acc::DeviceType deviceType) {
3951 return getValueInDeviceTypeSegment(getWorkerNumOperandsDeviceType(),
3952 getWorkerNumOperands(), deviceType);
3953}
3954
3955bool LoopOp::hasWorker() { return hasWorker(mlir::acc::DeviceType::None); }
3956
3957bool LoopOp::hasWorker(mlir::acc::DeviceType deviceType) {
3958 return hasDeviceType(getWorker(), deviceType);
3959}
3960
3961mlir::Operation::operand_range LoopOp::getTileValues() {
3962 return getTileValues(mlir::acc::DeviceType::None);
3963}
3964
3966LoopOp::getTileValues(mlir::acc::DeviceType deviceType) {
3967 return getValuesFromSegments(getTileOperandsDeviceType(), getTileOperands(),
3968 getTileOperandsSegments(), deviceType);
3969}
3970
3971std::optional<int64_t> LoopOp::getCollapseValue() {
3972 return getCollapseValue(mlir::acc::DeviceType::None);
3973}
3974
3975std::optional<int64_t>
3976LoopOp::getCollapseValue(mlir::acc::DeviceType deviceType) {
3977 if (!getCollapseAttr())
3978 return std::nullopt;
3979 if (auto pos = findSegment(getCollapseDeviceTypeAttr(), deviceType)) {
3980 auto intAttr =
3981 mlir::dyn_cast<IntegerAttr>(getCollapseAttr().getValue()[*pos]);
3982 return intAttr.getValue().getZExtValue();
3983 }
3984 return std::nullopt;
3985}
3986
3987mlir::Value LoopOp::getGangValue(mlir::acc::GangArgType gangArgType) {
3988 return getGangValue(gangArgType, mlir::acc::DeviceType::None);
3989}
3990
3991mlir::Value LoopOp::getGangValue(mlir::acc::GangArgType gangArgType,
3992 mlir::acc::DeviceType deviceType) {
3993 if (getGangOperands().empty())
3994 return {};
3995 if (auto pos = findSegment(*getGangOperandsDeviceType(), deviceType)) {
3996 int32_t nbOperandsBefore = 0;
3997 for (unsigned i = 0; i < *pos; ++i)
3998 nbOperandsBefore += (*getGangOperandsSegments())[i];
4000 getGangOperands()
4001 .drop_front(nbOperandsBefore)
4002 .take_front((*getGangOperandsSegments())[*pos]);
4003
4004 int32_t argTypeIdx = nbOperandsBefore;
4005 for (auto value : values) {
4006 auto gangArgTypeAttr = mlir::dyn_cast<mlir::acc::GangArgTypeAttr>(
4007 (*getGangOperandsArgType())[argTypeIdx]);
4008 if (gangArgTypeAttr.getValue() == gangArgType)
4009 return value;
4010 ++argTypeIdx;
4011 }
4012 }
4013 return {};
4014}
4015
4016bool LoopOp::hasGang() { return hasGang(mlir::acc::DeviceType::None); }
4017
4018bool LoopOp::hasGang(mlir::acc::DeviceType deviceType) {
4019 return hasDeviceType(getGang(), deviceType);
4020}
4021
4022llvm::SmallVector<mlir::Region *> acc::LoopOp::getLoopRegions() {
4023 return {&getRegion()};
4024}
4025
4026/// loop-control ::= `control` `(` ssa-id-and-type-list `)` `=`
4027/// `(` ssa-id-and-type-list `)` `to` `(` ssa-id-and-type-list `)` `step`
4028/// `(` ssa-id-and-type-list `)`
4029/// region
4030ParseResult
4033 SmallVectorImpl<Type> &lowerboundType,
4035 SmallVectorImpl<Type> &upperboundType,
4037 SmallVectorImpl<Type> &stepType) {
4038
4040 if (succeeded(
4041 parser.parseOptionalKeyword(acc::LoopOp::getControlKeyword()))) {
4042 if (parser.parseLParen() ||
4043 parser.parseArgumentList(inductionVars, OpAsmParser::Delimiter::None,
4044 /*allowType=*/true) ||
4045 parser.parseRParen() || parser.parseEqual() || parser.parseLParen() ||
4046 parser.parseOperandList(lowerbound, inductionVars.size(),
4048 parser.parseColonTypeList(lowerboundType) || parser.parseRParen() ||
4049 parser.parseKeyword("to") || parser.parseLParen() ||
4050 parser.parseOperandList(upperbound, inductionVars.size(),
4052 parser.parseColonTypeList(upperboundType) || parser.parseRParen() ||
4053 parser.parseKeyword("step") || parser.parseLParen() ||
4054 parser.parseOperandList(step, inductionVars.size(),
4056 parser.parseColonTypeList(stepType) || parser.parseRParen())
4057 return failure();
4058 }
4059 return parser.parseRegion(region, inductionVars);
4060}
4061
4063 ValueRange lowerbound, TypeRange lowerboundType,
4064 ValueRange upperbound, TypeRange upperboundType,
4065 ValueRange steps, TypeRange stepType) {
4066 ValueRange regionArgs = region.front().getArguments();
4067 if (!regionArgs.empty()) {
4068 p << acc::LoopOp::getControlKeyword() << "(";
4069 llvm::interleaveComma(regionArgs, p,
4070 [&p](Value v) { p << v << " : " << v.getType(); });
4071 p << ") = (" << lowerbound << " : " << lowerboundType << ") to ("
4072 << upperbound << " : " << upperboundType << ") " << " step (" << steps
4073 << " : " << stepType << ") ";
4074 }
4075 p.printRegion(region, /*printEntryBlockArgs=*/false);
4076}
4077
4078void acc::LoopOp::addSeq(MLIRContext *context,
4079 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4080 setSeqAttr(addDeviceTypeAffectedOperandHelper(context, getSeqAttr(),
4081 effectiveDeviceTypes));
4082}
4083
4084void acc::LoopOp::addIndependent(
4085 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4086 setIndependentAttr(addDeviceTypeAffectedOperandHelper(
4087 context, getIndependentAttr(), effectiveDeviceTypes));
4088}
4089
4090void acc::LoopOp::addAuto(MLIRContext *context,
4091 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4092 setAuto_Attr(addDeviceTypeAffectedOperandHelper(context, getAuto_Attr(),
4093 effectiveDeviceTypes));
4094}
4095
4096void acc::LoopOp::setCollapseForDeviceTypes(
4097 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
4098 llvm::APInt value) {
4101
4102 assert((getCollapseAttr() == nullptr) ==
4103 (getCollapseDeviceTypeAttr() == nullptr));
4104 assert(value.getBitWidth() == 64);
4105
4106 if (getCollapseAttr()) {
4107 for (const auto &existing :
4108 llvm::zip_equal(getCollapseAttr(), getCollapseDeviceTypeAttr())) {
4109 newValues.push_back(std::get<0>(existing));
4110 newDeviceTypes.push_back(std::get<1>(existing));
4111 }
4112 }
4113
4114 if (effectiveDeviceTypes.empty()) {
4115 // If the effective device-types list is empty, this is before there are any
4116 // being applied by device_type, so this should be added as a 'none'.
4117 newValues.push_back(
4118 mlir::IntegerAttr::get(mlir::IntegerType::get(context, 64), value));
4119 newDeviceTypes.push_back(
4120 acc::DeviceTypeAttr::get(context, DeviceType::None));
4121 } else {
4122 for (DeviceType dt : effectiveDeviceTypes) {
4123 newValues.push_back(
4124 mlir::IntegerAttr::get(mlir::IntegerType::get(context, 64), value));
4125 newDeviceTypes.push_back(acc::DeviceTypeAttr::get(context, dt));
4126 }
4127 }
4128
4129 setCollapseAttr(ArrayAttr::get(context, newValues));
4130 setCollapseDeviceTypeAttr(ArrayAttr::get(context, newDeviceTypes));
4131}
4132
4133void acc::LoopOp::setTileForDeviceTypes(
4134 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
4135 ValueRange values) {
4137 if (getTileOperandsSegments())
4138 llvm::copy(*getTileOperandsSegments(), std::back_inserter(segments));
4139
4140 setTileOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4141 context, getTileOperandsDeviceTypeAttr(), effectiveDeviceTypes, values,
4142 getTileOperandsMutable(), segments));
4143
4144 setTileOperandsSegments(segments);
4145}
4146
4147void acc::LoopOp::addVectorOperand(
4148 MLIRContext *context, mlir::Value newValue,
4149 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4150 setVectorOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4151 context, getVectorOperandsDeviceTypeAttr(), effectiveDeviceTypes,
4152 newValue, getVectorOperandsMutable()));
4153}
4154
4155void acc::LoopOp::addEmptyVector(
4156 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4157 setVectorAttr(addDeviceTypeAffectedOperandHelper(context, getVectorAttr(),
4158 effectiveDeviceTypes));
4159}
4160
4161void acc::LoopOp::addWorkerNumOperand(
4162 MLIRContext *context, mlir::Value newValue,
4163 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4164 setWorkerNumOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4165 context, getWorkerNumOperandsDeviceTypeAttr(), effectiveDeviceTypes,
4166 newValue, getWorkerNumOperandsMutable()));
4167}
4168
4169void acc::LoopOp::addEmptyWorker(
4170 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4171 setWorkerAttr(addDeviceTypeAffectedOperandHelper(context, getWorkerAttr(),
4172 effectiveDeviceTypes));
4173}
4174
4175void acc::LoopOp::addEmptyGang(
4176 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4177 setGangAttr(addDeviceTypeAffectedOperandHelper(context, getGangAttr(),
4178 effectiveDeviceTypes));
4179}
4180
4181bool acc::LoopOp::hasParallelismFlag(DeviceType dt) {
4182 auto hasDevice = [=](DeviceTypeAttr attr) -> bool {
4183 return attr.getValue() == dt;
4184 };
4185 auto testFromArr = [=](ArrayAttr arr) -> bool {
4186 return llvm::any_of(arr.getAsRange<DeviceTypeAttr>(), hasDevice);
4187 };
4188
4189 if (ArrayAttr arr = getSeqAttr(); arr && testFromArr(arr))
4190 return true;
4191 if (ArrayAttr arr = getIndependentAttr(); arr && testFromArr(arr))
4192 return true;
4193 if (ArrayAttr arr = getAuto_Attr(); arr && testFromArr(arr))
4194 return true;
4195
4196 return false;
4197}
4198
4199bool acc::LoopOp::hasDefaultGangWorkerVector() {
4200 return hasAnyGangWorkerVector(DeviceType::None);
4201}
4202
4203bool acc::LoopOp::hasAnyGangWorkerVector(DeviceType deviceType) {
4204 return hasVector(deviceType) || getVectorValue(deviceType) ||
4205 hasWorker(deviceType) || getWorkerValue(deviceType) ||
4206 hasGang(deviceType) || getGangValue(GangArgType::Num, deviceType) ||
4207 getGangValue(GangArgType::Dim, deviceType) ||
4208 getGangValue(GangArgType::Static, deviceType);
4209}
4210
4211acc::LoopParMode
4212acc::LoopOp::getDefaultOrDeviceTypeParallelism(DeviceType deviceType) {
4213 if (hasSeq(deviceType))
4214 return LoopParMode::loop_seq;
4215 if (hasAuto(deviceType))
4216 return LoopParMode::loop_auto;
4217 if (hasIndependent(deviceType))
4218 return LoopParMode::loop_independent;
4219 if (hasSeq())
4220 return LoopParMode::loop_seq;
4221 if (hasAuto())
4222 return LoopParMode::loop_auto;
4223 assert(hasIndependent() &&
4224 "loop must have default auto, seq, or independent");
4225 return LoopParMode::loop_independent;
4226}
4227
4228void acc::LoopOp::addGangOperands(
4229 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
4232 if (std::optional<ArrayRef<int32_t>> existingSegments =
4233 getGangOperandsSegments())
4234 llvm::copy(*existingSegments, std::back_inserter(segments));
4235
4236 unsigned beforeCount = segments.size();
4237
4238 setGangOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4239 context, getGangOperandsDeviceTypeAttr(), effectiveDeviceTypes, values,
4240 getGangOperandsMutable(), segments));
4241
4242 setGangOperandsSegments(segments);
4243
4244 // This is a bit of extra work to make sure we update the 'types' correctly by
4245 // adding to the types collection the correct number of times. We could
4246 // potentially add something similar to the
4247 // addDeviceTypeAffectedOperandHelper, but it seems that would be pretty
4248 // excessive for a one-off case.
4249 unsigned numAdded = segments.size() - beforeCount;
4250
4251 if (numAdded > 0) {
4253 if (getGangOperandsArgTypeAttr())
4254 llvm::copy(getGangOperandsArgTypeAttr(), std::back_inserter(gangTypes));
4255
4256 for (auto i : llvm::index_range(0u, numAdded)) {
4257 llvm::transform(argTypes, std::back_inserter(gangTypes),
4258 [=](mlir::acc::GangArgType gangTy) {
4259 return mlir::acc::GangArgTypeAttr::get(context, gangTy);
4260 });
4261 (void)i;
4262 }
4263
4264 setGangOperandsArgTypeAttr(mlir::ArrayAttr::get(context, gangTypes));
4265 }
4266}
4267
4268void acc::LoopOp::addPrivatization(MLIRContext *context,
4269 mlir::acc::PrivateOp op,
4270 mlir::acc::PrivateRecipeOp recipe) {
4271 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
4272 getPrivateOperandsMutable().append(op.getResult());
4273}
4274
4275void acc::LoopOp::addFirstPrivatization(
4276 MLIRContext *context, mlir::acc::FirstprivateOp op,
4277 mlir::acc::FirstprivateRecipeOp recipe) {
4278 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
4279 getFirstprivateOperandsMutable().append(op.getResult());
4280}
4281
4282void acc::LoopOp::addReduction(MLIRContext *context, mlir::acc::ReductionOp op,
4283 mlir::acc::ReductionRecipeOp recipe) {
4284 op.setRecipeAttr(mlir::SymbolRefAttr::get(context, recipe.getSymName()));
4285 getReductionOperandsMutable().append(op.getResult());
4286}
4287
4288//===----------------------------------------------------------------------===//
4289// DataOp
4290//===----------------------------------------------------------------------===//
4291
4292LogicalResult acc::DataOp::verify() {
4293 // 2.6.5. Data Construct restriction
4294 // At least one copy, copyin, copyout, create, no_create, present, deviceptr,
4295 // attach, or default clause must appear on a data construct.
4296 if (getOperands().empty() && !getDefaultAttr())
4297 return emitError("at least one operand or the default attribute "
4298 "must appear on the data operation");
4299
4300 for (mlir::Value operand : getDataClauseOperands())
4301 if (isa<BlockArgument>(operand) ||
4302 !mlir::isa<acc::AttachOp, acc::CopyinOp, acc::CopyoutOp, acc::CreateOp,
4303 acc::DeleteOp, acc::DetachOp, acc::DevicePtrOp,
4304 acc::GetDevicePtrOp, acc::NoCreateOp, acc::PresentOp,
4305 acc::MapInfoOp>(operand.getDefiningOp()))
4306 return emitError("expect data entry/exit operation or acc.getdeviceptr "
4307 "as defining op");
4308
4310 return failure();
4311
4312 return success();
4313}
4314
4315unsigned DataOp::getNumDataOperands() { return getDataClauseOperands().size(); }
4316
4317Value DataOp::getDataOperand(unsigned i) {
4318 unsigned numOptional = getIfCond() ? 1 : 0;
4319 numOptional += getAsyncOperands().size() ? 1 : 0;
4320 numOptional += getWaitOperands().size();
4321 return getOperand(numOptional + i);
4322}
4323
4324bool acc::DataOp::hasAsyncOnly() {
4325 return hasAsyncOnly(mlir::acc::DeviceType::None);
4326}
4327
4328bool acc::DataOp::hasAsyncOnly(mlir::acc::DeviceType deviceType) {
4329 return hasDeviceType(getAsyncOnly(), deviceType);
4330}
4331
4332mlir::Value DataOp::getAsyncValue() {
4333 return getAsyncValue(mlir::acc::DeviceType::None);
4334}
4335
4336mlir::Value DataOp::getAsyncValue(mlir::acc::DeviceType deviceType) {
4338 getAsyncOperands(), deviceType);
4339}
4340
4341bool DataOp::hasWaitOnly() { return hasWaitOnly(mlir::acc::DeviceType::None); }
4342
4343bool DataOp::hasWaitOnly(mlir::acc::DeviceType deviceType) {
4344 return hasDeviceType(getWaitOnly(), deviceType);
4345}
4346
4347mlir::Operation::operand_range DataOp::getWaitValues() {
4348 return getWaitValues(mlir::acc::DeviceType::None);
4349}
4350
4352DataOp::getWaitValues(mlir::acc::DeviceType deviceType) {
4354 getWaitOperandsDeviceType(), getWaitOperands(), getWaitOperandsSegments(),
4355 getHasWaitDevnum(), deviceType);
4356}
4357
4358mlir::Value DataOp::getWaitDevnum() {
4359 return getWaitDevnum(mlir::acc::DeviceType::None);
4360}
4361
4362mlir::Value DataOp::getWaitDevnum(mlir::acc::DeviceType deviceType) {
4363 return getWaitDevnumValue(getWaitOperandsDeviceType(), getWaitOperands(),
4364 getWaitOperandsSegments(), getHasWaitDevnum(),
4365 deviceType);
4366}
4367
4368void acc::DataOp::addAsyncOnly(
4369 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4370 setAsyncOnlyAttr(addDeviceTypeAffectedOperandHelper(
4371 context, getAsyncOnlyAttr(), effectiveDeviceTypes));
4372}
4373
4374void acc::DataOp::addAsyncOperand(
4375 MLIRContext *context, mlir::Value newValue,
4376 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4377 setAsyncOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4378 context, getAsyncOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValue,
4379 getAsyncOperandsMutable()));
4380}
4381
4382void acc::DataOp::addWaitOnly(MLIRContext *context,
4383 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4384 setWaitOnlyAttr(addDeviceTypeAffectedOperandHelper(context, getWaitOnlyAttr(),
4385 effectiveDeviceTypes));
4386}
4387
4388void acc::DataOp::addWaitOperands(
4389 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
4390 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4391
4393 if (getWaitOperandsSegments())
4394 llvm::copy(*getWaitOperandsSegments(), std::back_inserter(segments));
4395
4396 setWaitOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
4397 context, getWaitOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
4398 getWaitOperandsMutable(), segments));
4399 setWaitOperandsSegments(segments);
4400
4402 if (getHasWaitDevnumAttr())
4403 llvm::copy(getHasWaitDevnumAttr(), std::back_inserter(hasDevnums));
4404 hasDevnums.insert(
4405 hasDevnums.end(),
4406 std::max(effectiveDeviceTypes.size(), static_cast<size_t>(1)),
4407 mlir::BoolAttr::get(context, hasDevnum));
4408 setHasWaitDevnumAttr(mlir::ArrayAttr::get(context, hasDevnums));
4409}
4410
4411//===----------------------------------------------------------------------===//
4412// ExitDataOp
4413//===----------------------------------------------------------------------===//
4414
4415LogicalResult acc::ExitDataOp::verify() {
4416 // 2.6.6. Data Exit Directive restriction
4417 // At least one copyout, delete, or detach clause must appear on an exit data
4418 // directive.
4419 if (getDataClauseOperands().empty())
4420 return emitError("at least one operand must be present in dataOperands on "
4421 "the exit data operation");
4422
4423 // The async attribute represent the async clause without value. Therefore the
4424 // attribute and operand cannot appear at the same time.
4425 if (getAsyncOperand() && getAsync())
4426 return emitError("async attribute cannot appear with asyncOperand");
4427
4428 // The wait attribute represent the wait clause without values. Therefore the
4429 // attribute and operands cannot appear at the same time.
4430 if (!getWaitOperands().empty() && getWait())
4431 return emitError("wait attribute cannot appear with waitOperands");
4432
4433 if (getWaitDevnum() && getWaitOperands().empty())
4434 return emitError("wait_devnum cannot appear without waitOperands");
4435
4436 return success();
4437}
4438
4439unsigned ExitDataOp::getNumDataOperands() {
4440 return getDataClauseOperands().size();
4441}
4442
4443Value ExitDataOp::getDataOperand(unsigned i) {
4444 unsigned numOptional = getIfCond() ? 1 : 0;
4445 numOptional += getAsyncOperand() ? 1 : 0;
4446 numOptional += getWaitDevnum() ? 1 : 0;
4447 return getOperand(getWaitOperands().size() + numOptional + i);
4448}
4449
4450void ExitDataOp::getCanonicalizationPatterns(RewritePatternSet &results,
4451 MLIRContext *context) {
4452 results.add<RemoveConstantIfCondition<ExitDataOp>>(context);
4453}
4454
4455void ExitDataOp::addAsyncOnly(MLIRContext *context,
4456 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4457 assert(effectiveDeviceTypes.empty());
4458 assert(!getAsyncAttr());
4459 assert(!getAsyncOperand());
4460
4461 setAsyncAttr(mlir::UnitAttr::get(context));
4462}
4463
4464void ExitDataOp::addAsyncOperand(
4465 MLIRContext *context, mlir::Value newValue,
4466 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4467 assert(effectiveDeviceTypes.empty());
4468 assert(!getAsyncAttr());
4469 assert(!getAsyncOperand());
4470
4471 getAsyncOperandMutable().append(newValue);
4472}
4473
4474void ExitDataOp::addWaitOnly(MLIRContext *context,
4475 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4476 assert(effectiveDeviceTypes.empty());
4477
4478 if (getWaitAttr())
4479 return;
4480
4481 setWaitAttr(mlir::UnitAttr::get(context));
4482
4483 getWaitDevnumMutable().clear();
4484 getWaitOperandsMutable().clear();
4485}
4486
4487void ExitDataOp::addWaitOperands(
4488 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
4489 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4490 assert(effectiveDeviceTypes.empty());
4491
4492 if (getWaitAttr())
4493 return;
4494
4495 // FIXME: At one point we need to figure out how to support multiple devnums
4496 // here. For now, assert. Eventually we probably want to make dev-num and
4497 // operands work in 'lock-step', so that getWaitDevnum().size() ==
4498 // getWaitOperandsMutable().size().
4499 assert(!getWaitDevnum() && "Merging devnum not yet implemented");
4500
4501 // if hasDevnum, the first value is the devnum. The 'rest' go into the
4502 // operands list.
4503 if (hasDevnum) {
4504 getWaitDevnumMutable().append(newValues.front());
4505 newValues = newValues.drop_front();
4506 }
4507
4508 getWaitOperandsMutable().append(newValues);
4509}
4510
4511//===----------------------------------------------------------------------===//
4512// EnterDataOp
4513//===----------------------------------------------------------------------===//
4514
4515LogicalResult acc::EnterDataOp::verify() {
4516 // 2.6.6. Data Enter Directive restriction
4517 // At least one copyin, create, or attach clause must appear on an enter data
4518 // directive.
4519 if (getDataClauseOperands().empty())
4520 return emitError("at least one operand must be present in dataOperands on "
4521 "the enter data operation");
4522
4523 // The async attribute represent the async clause without value. Therefore the
4524 // attribute and operand cannot appear at the same time.
4525 if (getAsyncOperand() && getAsync())
4526 return emitError("async attribute cannot appear with asyncOperand");
4527
4528 // The wait attribute represent the wait clause without values. Therefore the
4529 // attribute and operands cannot appear at the same time.
4530 if (!getWaitOperands().empty() && getWait())
4531 return emitError("wait attribute cannot appear with waitOperands");
4532
4533 if (getWaitDevnum() && getWaitOperands().empty())
4534 return emitError("wait_devnum cannot appear without waitOperands");
4535
4536 for (mlir::Value operand : getDataClauseOperands())
4537 if (!mlir::isa<acc::AttachOp, acc::CreateOp, acc::CopyinOp, acc::MapInfoOp>(
4538 operand.getDefiningOp()))
4539 return emitError("expect data entry operation as defining op");
4540
4541 return success();
4542}
4543
4544unsigned EnterDataOp::getNumDataOperands() {
4545 return getDataClauseOperands().size();
4546}
4547
4548Value EnterDataOp::getDataOperand(unsigned i) {
4549 unsigned numOptional = getIfCond() ? 1 : 0;
4550 numOptional += getAsyncOperand() ? 1 : 0;
4551 numOptional += getWaitDevnum() ? 1 : 0;
4552 return getOperand(getWaitOperands().size() + numOptional + i);
4553}
4554
4555void EnterDataOp::getCanonicalizationPatterns(RewritePatternSet &results,
4556 MLIRContext *context) {
4557 results.add<RemoveConstantIfCondition<EnterDataOp>>(context);
4558}
4559
4560void EnterDataOp::addAsyncOnly(
4561 MLIRContext *context, llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4562 assert(effectiveDeviceTypes.empty());
4563 assert(!getAsyncAttr());
4564 assert(!getAsyncOperand());
4565
4566 setAsyncAttr(mlir::UnitAttr::get(context));
4567}
4568
4569void EnterDataOp::addAsyncOperand(
4570 MLIRContext *context, mlir::Value newValue,
4571 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4572 assert(effectiveDeviceTypes.empty());
4573 assert(!getAsyncAttr());
4574 assert(!getAsyncOperand());
4575
4576 getAsyncOperandMutable().append(newValue);
4577}
4578
4579void EnterDataOp::addWaitOnly(MLIRContext *context,
4580 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4581 assert(effectiveDeviceTypes.empty());
4582
4583 if (getWaitAttr())
4584 return;
4585
4586 setWaitAttr(mlir::UnitAttr::get(context));
4587
4588 getWaitDevnumMutable().clear();
4589 getWaitOperandsMutable().clear();
4590}
4591
4592void EnterDataOp::addWaitOperands(
4593 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
4594 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
4595 assert(effectiveDeviceTypes.empty());
4596
4597 if (getWaitAttr())
4598 return;
4599
4600 // FIXME: At one point we need to figure out how to support multiple devnums
4601 // here. For now, assert. Eventually we probably want to make dev-num and
4602 // operands work in 'lock-step', so that getWaitDevnum().size() ==
4603 // getWaitOperandsMutable().size().
4604 assert(!getWaitDevnum() && "Merging devnum not yet implemented");
4605
4606 // if hasDevnum, the first value is the devnum. The 'rest' go into the
4607 // operands list.
4608 if (hasDevnum) {
4609 getWaitDevnumMutable().append(newValues.front());
4610 newValues = newValues.drop_front();
4611 }
4612
4613 getWaitOperandsMutable().append(newValues);
4614}
4615
4616//===----------------------------------------------------------------------===//
4617// AtomicReadOp
4618//===----------------------------------------------------------------------===//
4619
4620LogicalResult AtomicReadOp::verify() { return verifyCommon(); }
4621
4622//===----------------------------------------------------------------------===//
4623// AtomicWriteOp
4624//===----------------------------------------------------------------------===//
4625
4626LogicalResult AtomicWriteOp::verify() { return verifyCommon(); }
4627
4628//===----------------------------------------------------------------------===//
4629// AtomicUpdateOp
4630//===----------------------------------------------------------------------===//
4631
4632LogicalResult AtomicUpdateOp::canonicalize(AtomicUpdateOp op,
4633 PatternRewriter &rewriter) {
4634 if (op.isNoOp()) {
4635 rewriter.eraseOp(op);
4636 return success();
4637 }
4638
4639 if (Value writeVal = op.getWriteOpVal()) {
4640 rewriter.replaceOpWithNewOp<AtomicWriteOp>(op, op.getX(), writeVal,
4641 op.getIfCond());
4642 return success();
4643 }
4644
4645 return failure();
4646}
4647
4648LogicalResult AtomicUpdateOp::verify() { return verifyCommon(); }
4649
4650LogicalResult AtomicUpdateOp::verifyRegions() { return verifyRegionsCommon(); }
4651
4652//===----------------------------------------------------------------------===//
4653// AtomicCaptureOp
4654//===----------------------------------------------------------------------===//
4655
4656AtomicReadOp AtomicCaptureOp::getAtomicReadOp() {
4657 if (auto op = dyn_cast<AtomicReadOp>(getFirstOp()))
4658 return op;
4659 return dyn_cast<AtomicReadOp>(getSecondOp());
4660}
4661
4662AtomicWriteOp AtomicCaptureOp::getAtomicWriteOp() {
4663 if (auto op = dyn_cast<AtomicWriteOp>(getFirstOp()))
4664 return op;
4665 return dyn_cast<AtomicWriteOp>(getSecondOp());
4666}
4667
4668AtomicUpdateOp AtomicCaptureOp::getAtomicUpdateOp() {
4669 if (auto op = dyn_cast<AtomicUpdateOp>(getFirstOp()))
4670 return op;
4671 return dyn_cast<AtomicUpdateOp>(getSecondOp());
4672}
4673
4674LogicalResult AtomicCaptureOp::verifyRegions() { return verifyRegionsCommon(); }
4675
4676//===----------------------------------------------------------------------===//
4677// DeclareEnterOp
4678//===----------------------------------------------------------------------===//
4679
4680template <typename Op>
4681static LogicalResult
4683 bool requireAtLeastOneOperand = true) {
4684 if (operands.empty() && requireAtLeastOneOperand)
4685 return emitError(
4686 op->getLoc(),
4687 "at least one operand must appear on the declare operation");
4688
4689 for (mlir::Value operand : operands) {
4690 if (isa<BlockArgument>(operand) ||
4691 !mlir::isa<acc::CopyinOp, acc::CopyoutOp, acc::CreateOp,
4692 acc::DevicePtrOp, acc::GetDevicePtrOp, acc::PresentOp,
4693 acc::DeclareDeviceResidentOp, acc::DeclareLinkOp,
4694 acc::MapInfoOp>(operand.getDefiningOp()))
4695 return op.emitError(
4696 "expect valid declare data entry operation or acc.getdeviceptr "
4697 "as defining op");
4698
4699 mlir::Value var{getVar(operand.getDefiningOp())};
4700 assert(var && "declare operands can only be data entry operations which "
4701 "must have var");
4702 (void)var;
4703 // acc.map_info encodes the clause effects in mapFlags instead.
4704 if (!mlir::isa<acc::MapInfoOp>(operand.getDefiningOp())) {
4705 std::optional<mlir::acc::DataClause> dataClauseOptional{
4706 getDataClause(operand.getDefiningOp())};
4707 assert(dataClauseOptional.has_value() &&
4708 "declare operands can only be data entry operations which must "
4709 "have dataClause");
4710 (void)dataClauseOptional;
4711 }
4712 }
4713
4714 return success();
4715}
4716
4717LogicalResult acc::DeclareEnterOp::verify() {
4718 return checkDeclareOperands(*this, this->getDataClauseOperands());
4719}
4720
4721//===----------------------------------------------------------------------===//
4722// DeclareExitOp
4723//===----------------------------------------------------------------------===//
4724
4725LogicalResult acc::DeclareExitOp::verify() {
4726 if (getToken())
4727 return checkDeclareOperands(*this, this->getDataClauseOperands(),
4728 /*requireAtLeastOneOperand=*/false);
4729 return checkDeclareOperands(*this, this->getDataClauseOperands());
4730}
4731
4732//===----------------------------------------------------------------------===//
4733// DeclareOp
4734//===----------------------------------------------------------------------===//
4735
4736LogicalResult acc::DeclareOp::verify() {
4737 return checkDeclareOperands(*this, this->getDataClauseOperands());
4738}
4739
4740//===----------------------------------------------------------------------===//
4741// RoutineOp
4742//===----------------------------------------------------------------------===//
4743
4744static unsigned getParallelismForDeviceType(acc::RoutineOp op,
4745 acc::DeviceType dtype) {
4746 unsigned parallelism = 0;
4747 parallelism += (op.hasGang(dtype) || op.getGangDimValue(dtype)) ? 1 : 0;
4748 parallelism += op.hasWorker(dtype) ? 1 : 0;
4749 parallelism += op.hasVector(dtype) ? 1 : 0;
4750 parallelism += op.hasSeq(dtype) ? 1 : 0;
4751 return parallelism;
4752}
4753
4754LogicalResult acc::RoutineOp::verify() {
4755 unsigned baseParallelism =
4756 getParallelismForDeviceType(*this, acc::DeviceType::None);
4757
4758 if (baseParallelism > 1)
4759 return emitError() << "only one of `gang`, `worker`, `vector`, `seq` can "
4760 "be present at the same time";
4761
4762 for (uint32_t dtypeInt = 0; dtypeInt != acc::getMaxEnumValForDeviceType();
4763 ++dtypeInt) {
4764 auto dtype = static_cast<acc::DeviceType>(dtypeInt);
4765 if (dtype == acc::DeviceType::None)
4766 continue;
4767 unsigned parallelism = getParallelismForDeviceType(*this, dtype);
4768
4769 if (parallelism > 1 || (baseParallelism == 1 && parallelism == 1))
4770 return emitError() << "only one of `gang`, `worker`, `vector`, `seq` can "
4771 "be present at the same time for device_type `"
4772 << acc::stringifyDeviceType(dtype) << "`";
4773 }
4774
4775 return success();
4776}
4777
4778static ParseResult parseBindName(OpAsmParser &parser,
4779 mlir::ArrayAttr &bindIdName,
4780 mlir::ArrayAttr &bindStrName,
4781 mlir::ArrayAttr &deviceIdTypes,
4782 mlir::ArrayAttr &deviceStrTypes) {
4783 llvm::SmallVector<mlir::Attribute> bindIdNameAttrs;
4784 llvm::SmallVector<mlir::Attribute> bindStrNameAttrs;
4785 llvm::SmallVector<mlir::Attribute> deviceIdTypeAttrs;
4786 llvm::SmallVector<mlir::Attribute> deviceStrTypeAttrs;
4787
4788 if (failed(parser.parseCommaSeparatedList([&]() {
4789 llvm::SMLoc attrLoc = parser.getCurrentLocation();
4790 mlir::Attribute newAttr;
4791 bool isSymbolRefAttr;
4792 if (parser.parseAttribute(newAttr))
4793 return failure();
4794 if (auto symbolRefAttr = dyn_cast<mlir::SymbolRefAttr>(newAttr)) {
4795 bindIdNameAttrs.push_back(symbolRefAttr);
4796 isSymbolRefAttr = true;
4797 } else if (auto stringAttr = dyn_cast<mlir::StringAttr>(newAttr)) {
4798 bindStrNameAttrs.push_back(stringAttr);
4799 isSymbolRefAttr = false;
4800 } else {
4801 parser.emitError(attrLoc,
4802 "expected symbol reference or string attribute");
4803 return failure();
4804 }
4805 if (failed(parser.parseOptionalLSquare())) {
4806 if (isSymbolRefAttr) {
4807 deviceIdTypeAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
4808 parser.getContext(), mlir::acc::DeviceType::None));
4809 } else {
4810 deviceStrTypeAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
4811 parser.getContext(), mlir::acc::DeviceType::None));
4812 }
4813 } else {
4814 if (isSymbolRefAttr) {
4815 if (parser.parseAttribute(deviceIdTypeAttrs.emplace_back()) ||
4816 parser.parseRSquare())
4817 return failure();
4818 } else {
4819 if (parser.parseAttribute(deviceStrTypeAttrs.emplace_back()) ||
4820 parser.parseRSquare())
4821 return failure();
4822 }
4823 }
4824 return success();
4825 })))
4826 return failure();
4827
4828 bindIdName = ArrayAttr::get(parser.getContext(), bindIdNameAttrs);
4829 bindStrName = ArrayAttr::get(parser.getContext(), bindStrNameAttrs);
4830 deviceIdTypes = ArrayAttr::get(parser.getContext(), deviceIdTypeAttrs);
4831 deviceStrTypes = ArrayAttr::get(parser.getContext(), deviceStrTypeAttrs);
4832
4833 return success();
4834}
4835
4837 std::optional<mlir::ArrayAttr> bindIdName,
4838 std::optional<mlir::ArrayAttr> bindStrName,
4839 std::optional<mlir::ArrayAttr> deviceIdTypes,
4840 std::optional<mlir::ArrayAttr> deviceStrTypes) {
4841 // Create combined vectors for all bind names and device types
4844
4845 // Append bindIdName and deviceIdTypes
4846 if (hasDeviceTypeValues(deviceIdTypes)) {
4847 allBindNames.append(bindIdName->begin(), bindIdName->end());
4848 allDeviceTypes.append(deviceIdTypes->begin(), deviceIdTypes->end());
4849 }
4850
4851 // Append bindStrName and deviceStrTypes
4852 if (hasDeviceTypeValues(deviceStrTypes)) {
4853 allBindNames.append(bindStrName->begin(), bindStrName->end());
4854 allDeviceTypes.append(deviceStrTypes->begin(), deviceStrTypes->end());
4855 }
4856
4857 // Print the combined sequence
4858 if (!allBindNames.empty())
4859 llvm::interleaveComma(llvm::zip(allBindNames, allDeviceTypes), p,
4860 [&](const auto &pair) {
4861 p << std::get<0>(pair);
4862 printSingleDeviceType(p, std::get<1>(pair));
4863 });
4864}
4865
4866static ParseResult parseRoutineGangClause(OpAsmParser &parser,
4867 mlir::ArrayAttr &gang,
4868 mlir::ArrayAttr &gangDim,
4869 mlir::ArrayAttr &gangDimDeviceTypes) {
4870
4871 llvm::SmallVector<mlir::Attribute> gangAttrs, gangDimAttrs,
4872 gangDimDeviceTypeAttrs;
4873 bool needCommaBeforeOperands = false;
4874
4875 // Gang keyword only
4876 if (failed(parser.parseOptionalLParen())) {
4877 gangAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
4878 parser.getContext(), mlir::acc::DeviceType::None));
4879 gang = ArrayAttr::get(parser.getContext(), gangAttrs);
4880 return success();
4881 }
4882
4883 // Parse keyword only attributes
4884 if (succeeded(parser.parseOptionalLSquare())) {
4885 if (failed(parser.parseCommaSeparatedList([&]() {
4886 if (parser.parseAttribute(gangAttrs.emplace_back()))
4887 return failure();
4888 return success();
4889 })))
4890 return failure();
4891 if (parser.parseRSquare())
4892 return failure();
4893 needCommaBeforeOperands = true;
4894 }
4895
4896 if (needCommaBeforeOperands && failed(parser.parseComma()))
4897 return failure();
4898
4899 if (failed(parser.parseCommaSeparatedList([&]() {
4900 if (parser.parseKeyword(acc::RoutineOp::getGangDimKeyword()) ||
4901 parser.parseColon() ||
4902 parser.parseAttribute(gangDimAttrs.emplace_back()))
4903 return failure();
4904 if (succeeded(parser.parseOptionalLSquare())) {
4905 if (parser.parseAttribute(gangDimDeviceTypeAttrs.emplace_back()) ||
4906 parser.parseRSquare())
4907 return failure();
4908 } else {
4909 gangDimDeviceTypeAttrs.push_back(mlir::acc::DeviceTypeAttr::get(
4910 parser.getContext(), mlir::acc::DeviceType::None));
4911 }
4912 return success();
4913 })))
4914 return failure();
4915
4916 if (failed(parser.parseRParen()))
4917 return failure();
4918
4919 gang = ArrayAttr::get(parser.getContext(), gangAttrs);
4920 gangDim = ArrayAttr::get(parser.getContext(), gangDimAttrs);
4921 gangDimDeviceTypes =
4922 ArrayAttr::get(parser.getContext(), gangDimDeviceTypeAttrs);
4923
4924 return success();
4925}
4926
4928 std::optional<mlir::ArrayAttr> gang,
4929 std::optional<mlir::ArrayAttr> gangDim,
4930 std::optional<mlir::ArrayAttr> gangDimDeviceTypes) {
4931
4932 if (!hasDeviceTypeValues(gangDimDeviceTypes) && hasDeviceTypeValues(gang) &&
4933 gang->size() == 1) {
4934 auto deviceTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>((*gang)[0]);
4935 if (deviceTypeAttr.getValue() == mlir::acc::DeviceType::None)
4936 return;
4937 }
4938
4939 p << "(";
4940
4941 printDeviceTypes(p, gang);
4942
4943 if (hasDeviceTypeValues(gang) && hasDeviceTypeValues(gangDimDeviceTypes))
4944 p << ", ";
4945
4946 if (hasDeviceTypeValues(gangDimDeviceTypes))
4947 llvm::interleaveComma(llvm::zip(*gangDim, *gangDimDeviceTypes), p,
4948 [&](const auto &pair) {
4949 p << acc::RoutineOp::getGangDimKeyword() << ": ";
4950 p << std::get<0>(pair);
4951 printSingleDeviceType(p, std::get<1>(pair));
4952 });
4953
4954 p << ")";
4955}
4956
4957static ParseResult parseDeviceTypeArrayAttr(OpAsmParser &parser,
4958 mlir::ArrayAttr &deviceTypes) {
4960 // Keyword only
4961 if (failed(parser.parseOptionalLParen())) {
4962 attributes.push_back(mlir::acc::DeviceTypeAttr::get(
4963 parser.getContext(), mlir::acc::DeviceType::None));
4964 deviceTypes = ArrayAttr::get(parser.getContext(), attributes);
4965 return success();
4966 }
4967
4968 // Parse device type attributes
4969 if (succeeded(parser.parseOptionalLSquare())) {
4970 if (failed(parser.parseCommaSeparatedList([&]() {
4971 if (parser.parseAttribute(attributes.emplace_back()))
4972 return failure();
4973 return success();
4974 })))
4975 return failure();
4976 if (parser.parseRSquare() || parser.parseRParen())
4977 return failure();
4978 }
4979 deviceTypes = ArrayAttr::get(parser.getContext(), attributes);
4980 return success();
4981}
4982
4983static void
4985 std::optional<mlir::ArrayAttr> deviceTypes) {
4986
4987 if (hasDeviceTypeValues(deviceTypes) && deviceTypes->size() == 1) {
4988 auto deviceTypeAttr =
4989 mlir::dyn_cast<mlir::acc::DeviceTypeAttr>((*deviceTypes)[0]);
4990 if (deviceTypeAttr.getValue() == mlir::acc::DeviceType::None)
4991 return;
4992 }
4993
4994 if (!hasDeviceTypeValues(deviceTypes))
4995 return;
4996
4997 p << "([";
4998 llvm::interleaveComma(*deviceTypes, p, [&](mlir::Attribute attr) {
4999 auto dTypeAttr = mlir::dyn_cast<mlir::acc::DeviceTypeAttr>(attr);
5000 p << dTypeAttr;
5001 });
5002 p << "])";
5003}
5004
5005bool RoutineOp::hasWorker() { return hasWorker(mlir::acc::DeviceType::None); }
5006
5007bool RoutineOp::hasWorker(mlir::acc::DeviceType deviceType) {
5008 return hasDeviceType(getWorker(), deviceType);
5009}
5010
5011bool RoutineOp::hasVector() { return hasVector(mlir::acc::DeviceType::None); }
5012
5013bool RoutineOp::hasVector(mlir::acc::DeviceType deviceType) {
5014 return hasDeviceType(getVector(), deviceType);
5015}
5016
5017bool RoutineOp::hasSeq() { return hasSeq(mlir::acc::DeviceType::None); }
5018
5019bool RoutineOp::hasSeq(mlir::acc::DeviceType deviceType) {
5020 return hasDeviceType(getSeq(), deviceType);
5021}
5022
5023std::optional<std::variant<mlir::SymbolRefAttr, mlir::StringAttr>>
5024RoutineOp::getBindNameValue() {
5025 return getBindNameValue(mlir::acc::DeviceType::None);
5026}
5027
5028std::optional<std::variant<mlir::SymbolRefAttr, mlir::StringAttr>>
5029RoutineOp::getBindNameValue(mlir::acc::DeviceType deviceType) {
5030 if (hasDeviceTypeValues(getBindIdNameDeviceType())) {
5031 if (auto pos = findSegment(*getBindIdNameDeviceType(), deviceType)) {
5032 auto attr = (*getBindIdName())[*pos];
5033 auto symbolRefAttr = dyn_cast<mlir::SymbolRefAttr>(attr);
5034 assert(symbolRefAttr && "expected SymbolRef");
5035 return symbolRefAttr;
5036 }
5037 }
5038
5039 if (hasDeviceTypeValues(getBindStrNameDeviceType())) {
5040 if (auto pos = findSegment(*getBindStrNameDeviceType(), deviceType)) {
5041 auto attr = (*getBindStrName())[*pos];
5042 auto stringAttr = dyn_cast<mlir::StringAttr>(attr);
5043 assert(stringAttr && "expected String");
5044 return stringAttr;
5045 }
5046 }
5047
5048 return std::nullopt;
5049}
5050
5051bool RoutineOp::hasGang() { return hasGang(mlir::acc::DeviceType::None); }
5052
5053bool RoutineOp::hasGang(mlir::acc::DeviceType deviceType) {
5054 return hasDeviceType(getGang(), deviceType);
5055}
5056
5057std::optional<int64_t> RoutineOp::getGangDimValue() {
5058 return getGangDimValue(mlir::acc::DeviceType::None);
5059}
5060
5061std::optional<int64_t>
5062RoutineOp::getGangDimValue(mlir::acc::DeviceType deviceType) {
5063 if (!hasDeviceTypeValues(getGangDimDeviceType()))
5064 return std::nullopt;
5065 if (auto pos = findSegment(*getGangDimDeviceType(), deviceType)) {
5066 auto intAttr = mlir::dyn_cast<mlir::IntegerAttr>((*getGangDim())[*pos]);
5067 return intAttr.getInt();
5068 }
5069 return std::nullopt;
5070}
5071
5072void RoutineOp::addSeq(MLIRContext *context,
5073 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5074 setSeqAttr(addDeviceTypeAffectedOperandHelper(context, getSeqAttr(),
5075 effectiveDeviceTypes));
5076}
5077
5078void RoutineOp::addVector(MLIRContext *context,
5079 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5080 setVectorAttr(addDeviceTypeAffectedOperandHelper(context, getVectorAttr(),
5081 effectiveDeviceTypes));
5082}
5083
5084void RoutineOp::addWorker(MLIRContext *context,
5085 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5086 setWorkerAttr(addDeviceTypeAffectedOperandHelper(context, getWorkerAttr(),
5087 effectiveDeviceTypes));
5088}
5089
5090void RoutineOp::addGang(MLIRContext *context,
5091 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5092 setGangAttr(addDeviceTypeAffectedOperandHelper(context, getGangAttr(),
5093 effectiveDeviceTypes));
5094}
5095
5096void RoutineOp::addGang(MLIRContext *context,
5097 llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
5098 uint64_t val) {
5101
5102 if (getGangDimAttr())
5103 llvm::copy(getGangDimAttr(), std::back_inserter(dimValues));
5104 if (getGangDimDeviceTypeAttr())
5105 llvm::copy(getGangDimDeviceTypeAttr(), std::back_inserter(deviceTypes));
5106
5107 assert(dimValues.size() == deviceTypes.size());
5108
5109 if (effectiveDeviceTypes.empty()) {
5110 dimValues.push_back(
5111 mlir::IntegerAttr::get(mlir::IntegerType::get(context, 64), val));
5112 deviceTypes.push_back(
5113 acc::DeviceTypeAttr::get(context, acc::DeviceType::None));
5114 } else {
5115 for (DeviceType dt : effectiveDeviceTypes) {
5116 dimValues.push_back(
5117 mlir::IntegerAttr::get(mlir::IntegerType::get(context, 64), val));
5118 deviceTypes.push_back(acc::DeviceTypeAttr::get(context, dt));
5119 }
5120 }
5121 assert(dimValues.size() == deviceTypes.size());
5122
5123 setGangDimAttr(mlir::ArrayAttr::get(context, dimValues));
5124 setGangDimDeviceTypeAttr(mlir::ArrayAttr::get(context, deviceTypes));
5125}
5126
5127void RoutineOp::addBindStrName(MLIRContext *context,
5128 llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
5129 mlir::StringAttr val) {
5130 unsigned before = getBindStrNameDeviceTypeAttr()
5131 ? getBindStrNameDeviceTypeAttr().size()
5132 : 0;
5133
5134 setBindStrNameDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
5135 context, getBindStrNameDeviceTypeAttr(), effectiveDeviceTypes));
5136 unsigned after = getBindStrNameDeviceTypeAttr().size();
5137
5139 if (getBindStrNameAttr())
5140 llvm::copy(getBindStrNameAttr(), std::back_inserter(vals));
5141 for (unsigned i = 0; i < after - before; ++i)
5142 vals.push_back(val);
5143
5144 setBindStrNameAttr(mlir::ArrayAttr::get(context, vals));
5145}
5146
5147void RoutineOp::addBindIDName(MLIRContext *context,
5148 llvm::ArrayRef<DeviceType> effectiveDeviceTypes,
5149 mlir::SymbolRefAttr val) {
5150 unsigned before =
5151 getBindIdNameDeviceTypeAttr() ? getBindIdNameDeviceTypeAttr().size() : 0;
5152
5153 setBindIdNameDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
5154 context, getBindIdNameDeviceTypeAttr(), effectiveDeviceTypes));
5155 unsigned after = getBindIdNameDeviceTypeAttr().size();
5156
5158 if (getBindIdNameAttr())
5159 llvm::copy(getBindIdNameAttr(), std::back_inserter(vals));
5160 for (unsigned i = 0; i < after - before; ++i)
5161 vals.push_back(val);
5162
5163 setBindIdNameAttr(mlir::ArrayAttr::get(context, vals));
5164}
5165
5166//===----------------------------------------------------------------------===//
5167// InitOp
5168//===----------------------------------------------------------------------===//
5169
5170LogicalResult acc::InitOp::verify() {
5171 if (getOperation()->getParentOfType<ACC_COMPUTE_CONSTRUCT_AND_LOOP_OPS>())
5172 return emitOpError("cannot be nested in a compute operation");
5173 return success();
5174}
5175
5176void acc::InitOp::addDeviceType(MLIRContext *context,
5177 mlir::acc::DeviceType deviceType) {
5179 if (getDeviceTypesAttr())
5180 llvm::copy(getDeviceTypesAttr(), std::back_inserter(deviceTypes));
5181
5182 deviceTypes.push_back(acc::DeviceTypeAttr::get(context, deviceType));
5183 setDeviceTypesAttr(mlir::ArrayAttr::get(context, deviceTypes));
5184}
5185
5186//===----------------------------------------------------------------------===//
5187// ShutdownOp
5188//===----------------------------------------------------------------------===//
5189
5190LogicalResult acc::ShutdownOp::verify() {
5191 if (getOperation()->getParentOfType<ACC_COMPUTE_CONSTRUCT_AND_LOOP_OPS>())
5192 return emitOpError("cannot be nested in a compute operation");
5193 return success();
5194}
5195
5196void acc::ShutdownOp::addDeviceType(MLIRContext *context,
5197 mlir::acc::DeviceType deviceType) {
5199 if (getDeviceTypesAttr())
5200 llvm::copy(getDeviceTypesAttr(), std::back_inserter(deviceTypes));
5201
5202 deviceTypes.push_back(acc::DeviceTypeAttr::get(context, deviceType));
5203 setDeviceTypesAttr(mlir::ArrayAttr::get(context, deviceTypes));
5204}
5205
5206//===----------------------------------------------------------------------===//
5207// SetOp
5208//===----------------------------------------------------------------------===//
5209
5210LogicalResult acc::SetOp::verify() {
5211 if (getOperation()->getParentOfType<ACC_COMPUTE_CONSTRUCT_AND_LOOP_OPS>())
5212 return emitOpError("cannot be nested in a compute operation");
5213 if (!getDeviceTypeAttr() && !getDefaultAsync() && !getDeviceNum())
5214 return emitOpError("at least one default_async, device_num, or device_type "
5215 "operand must appear");
5216 return success();
5217}
5218
5219//===----------------------------------------------------------------------===//
5220// UpdateOp
5221//===----------------------------------------------------------------------===//
5222
5223LogicalResult acc::UpdateOp::verify() {
5224 // At least one of host or device should have a value.
5225 if (getDataClauseOperands().empty())
5226 return emitError("at least one value must be present in dataOperands");
5227
5229 getAsyncOperandsDeviceTypeAttr(),
5230 "async")))
5231 return failure();
5232
5234 *this, getWaitOperands(), getWaitOperandsSegmentsAttr(),
5235 getWaitOperandsDeviceTypeAttr(), "wait")))
5236 return failure();
5237
5239 return failure();
5240
5241 for (mlir::Value operand : getDataClauseOperands())
5242 if (!mlir::isa<acc::UpdateDeviceOp, acc::UpdateHostOp, acc::GetDevicePtrOp,
5243 acc::MapInfoOp>(operand.getDefiningOp()))
5244 return emitError("expect data entry/exit operation or acc.getdeviceptr "
5245 "as defining op");
5246
5247 return success();
5248}
5249
5250unsigned UpdateOp::getNumDataOperands() {
5251 return getDataClauseOperands().size();
5252}
5253
5254Value UpdateOp::getDataOperand(unsigned i) {
5255 unsigned numOptional = getAsyncOperands().size();
5256 numOptional += getIfCond() ? 1 : 0;
5257 return getOperand(getWaitOperands().size() + numOptional + i);
5258}
5259
5260void UpdateOp::getCanonicalizationPatterns(RewritePatternSet &results,
5261 MLIRContext *context) {
5262 results.add<RemoveConstantIfCondition<UpdateOp>>(context);
5263}
5264
5265bool UpdateOp::hasAsyncOnly() {
5266 return hasAsyncOnly(mlir::acc::DeviceType::None);
5267}
5268
5269bool UpdateOp::hasAsyncOnly(mlir::acc::DeviceType deviceType) {
5270 return hasDeviceType(getAsyncOnly(), deviceType);
5271}
5272
5273mlir::Value UpdateOp::getAsyncValue() {
5274 return getAsyncValue(mlir::acc::DeviceType::None);
5275}
5276
5277mlir::Value UpdateOp::getAsyncValue(mlir::acc::DeviceType deviceType) {
5279 return {};
5280
5281 if (auto pos = findSegment(*getAsyncOperandsDeviceType(), deviceType))
5282 return getAsyncOperands()[*pos];
5283
5284 return {};
5285}
5286
5287bool UpdateOp::hasWaitOnly() {
5288 return hasWaitOnly(mlir::acc::DeviceType::None);
5289}
5290
5291bool UpdateOp::hasWaitOnly(mlir::acc::DeviceType deviceType) {
5292 return hasDeviceType(getWaitOnly(), deviceType);
5293}
5294
5295mlir::Operation::operand_range UpdateOp::getWaitValues() {
5296 return getWaitValues(mlir::acc::DeviceType::None);
5297}
5298
5300UpdateOp::getWaitValues(mlir::acc::DeviceType deviceType) {
5302 getWaitOperandsDeviceType(), getWaitOperands(), getWaitOperandsSegments(),
5303 getHasWaitDevnum(), deviceType);
5304}
5305
5306mlir::Value UpdateOp::getWaitDevnum() {
5307 return getWaitDevnum(mlir::acc::DeviceType::None);
5308}
5309
5310mlir::Value UpdateOp::getWaitDevnum(mlir::acc::DeviceType deviceType) {
5311 return getWaitDevnumValue(getWaitOperandsDeviceType(), getWaitOperands(),
5312 getWaitOperandsSegments(), getHasWaitDevnum(),
5313 deviceType);
5314}
5315
5316void UpdateOp::addAsyncOnly(MLIRContext *context,
5317 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5318 setAsyncOnlyAttr(addDeviceTypeAffectedOperandHelper(
5319 context, getAsyncOnlyAttr(), effectiveDeviceTypes));
5320}
5321
5322void UpdateOp::addAsyncOperand(
5323 MLIRContext *context, mlir::Value newValue,
5324 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5325 setAsyncOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
5326 context, getAsyncOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValue,
5327 getAsyncOperandsMutable()));
5328}
5329
5330void UpdateOp::addWaitOnly(MLIRContext *context,
5331 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5332 setWaitOnlyAttr(addDeviceTypeAffectedOperandHelper(context, getWaitOnlyAttr(),
5333 effectiveDeviceTypes));
5334}
5335
5336void UpdateOp::addWaitOperands(
5337 MLIRContext *context, bool hasDevnum, mlir::ValueRange newValues,
5338 llvm::ArrayRef<DeviceType> effectiveDeviceTypes) {
5339
5341 if (getWaitOperandsSegments())
5342 llvm::copy(*getWaitOperandsSegments(), std::back_inserter(segments));
5343
5344 setWaitOperandsDeviceTypeAttr(addDeviceTypeAffectedOperandHelper(
5345 context, getWaitOperandsDeviceTypeAttr(), effectiveDeviceTypes, newValues,
5346 getWaitOperandsMutable(), segments));
5347 setWaitOperandsSegments(segments);
5348
5350 if (getHasWaitDevnumAttr())
5351 llvm::copy(getHasWaitDevnumAttr(), std::back_inserter(hasDevnums));
5352 hasDevnums.insert(
5353 hasDevnums.end(),
5354 std::max(effectiveDeviceTypes.size(), static_cast<size_t>(1)),
5355 mlir::BoolAttr::get(context, hasDevnum));
5356 setHasWaitDevnumAttr(mlir::ArrayAttr::get(context, hasDevnums));
5357}
5358
5359//===----------------------------------------------------------------------===//
5360// WaitOp
5361//===----------------------------------------------------------------------===//
5362
5363LogicalResult acc::WaitOp::verify() {
5364 // The async attribute represent the async clause without value. Therefore the
5365 // attribute and operand cannot appear at the same time.
5366 if (getAsyncOperand() && getAsync())
5367 return emitError("async attribute cannot appear with asyncOperand");
5368
5369 if (getWaitDevnum() && getWaitOperands().empty())
5370 return emitError("wait_devnum cannot appear without waitOperands");
5371
5372 return success();
5373}
5374
5375#define GET_OP_CLASSES
5376#include "mlir/Dialect/OpenACC/OpenACCOps.cpp.inc"
5377
5378#define GET_ATTRDEF_CLASSES
5379#include "mlir/Dialect/OpenACC/OpenACCOpsAttributes.cpp.inc"
5380
5381#define GET_TYPEDEF_CLASSES
5382#include "mlir/Dialect/OpenACC/OpenACCOpsTypes.cpp.inc"
5383
5384//===----------------------------------------------------------------------===//
5385// acc dialect utilities
5386//===----------------------------------------------------------------------===//
5387
5390 auto varPtr{llvm::TypeSwitch<mlir::Operation *,
5392 accDataClauseOp)
5393 .Case<ACC_DATA_ENTRY_OPS, mlir::acc::MapInfoOp>(
5394 [&](auto entry) { return entry.getVarPtr(); })
5395 .Case<mlir::acc::CopyoutOp, mlir::acc::UpdateHostOp>(
5396 [&](auto exit) { return exit.getVarPtr(); })
5397 .Default([&](mlir::Operation *) {
5399 })};
5400 return varPtr;
5401}
5402
5404 auto varPtr{llvm::TypeSwitch<mlir::Operation *, mlir::Value>(accDataClauseOp)
5405 .Case<ACC_DATA_ENTRY_OPS, mlir::acc::MapInfoOp>(
5406 [&](auto entry) { return entry.getVar(); })
5407 .Default([&](mlir::Operation *) { return mlir::Value(); })};
5408 return varPtr;
5409}
5410
5412 auto varType{llvm::TypeSwitch<mlir::Operation *, mlir::Type>(accDataClauseOp)
5413 .Case<ACC_DATA_ENTRY_OPS, mlir::acc::MapInfoOp>(
5414 [&](auto entry) { return entry.getVarType(); })
5415 .Case<mlir::acc::CopyoutOp, mlir::acc::UpdateHostOp>(
5416 [&](auto exit) { return exit.getVarType(); })
5417 .Default([&](mlir::Operation *) { return mlir::Type(); })};
5418 return varType;
5419}
5420
5423 auto accPtr{
5426 accDataClauseOp)
5427 .Case<ACC_DATA_ENTRY_OPS, ACC_DATA_EXIT_OPS, mlir::acc::MapInfoOp>(
5428 [&](auto dataClause) { return dataClause.getAccPtr(); })
5429 .Default([&](mlir::Operation *) {
5431 })};
5432 return accPtr;
5433}
5434
5436 auto accPtr{
5438 .Case<ACC_DATA_ENTRY_OPS, ACC_DATA_EXIT_OPS, mlir::acc::MapInfoOp>(
5439 [&](auto dataClause) { return dataClause.getAccVar(); })
5440 .Default([&](mlir::Operation *) { return mlir::Value(); })};
5441 return accPtr;
5442}
5443
5445 auto varPtrPtr{
5447 .Case<ACC_DATA_ENTRY_OPS, mlir::acc::MapInfoOp>(
5448 [&](auto dataClause) { return dataClause.getVarPtrPtr(); })
5449 .Default([&](mlir::Operation *) { return mlir::Value(); })};
5450 return varPtrPtr;
5451}
5452
5457 accDataClauseOp)
5458 .Case<ACC_DATA_ENTRY_OPS, ACC_DATA_EXIT_OPS, mlir::acc::MapInfoOp>(
5459 [&](auto dataClause) {
5461 dataClause.getBounds().begin(),
5462 dataClause.getBounds().end());
5463 })
5464 .Default([&](mlir::Operation *) {
5466 })};
5467 return bounds;
5468}
5469
5473 accDataClauseOp)
5474 .Case<ACC_DATA_ENTRY_OPS, ACC_DATA_EXIT_OPS>([&](auto dataClause) {
5476 dataClause.getAsyncOperands().begin(),
5477 dataClause.getAsyncOperands().end());
5478 })
5479 .Default([&](mlir::Operation *) {
5481 });
5482}
5483
5484mlir::ArrayAttr
5487 .Case<ACC_DATA_ENTRY_OPS, ACC_DATA_EXIT_OPS>([&](auto dataClause) {
5488 return dataClause.getAsyncOperandsDeviceTypeAttr();
5489 })
5490 .Default([&](mlir::Operation *) { return mlir::ArrayAttr{}; });
5491}
5492
5493mlir::ArrayAttr mlir::acc::getAsyncOnly(mlir::Operation *accDataClauseOp) {
5496 [&](auto dataClause) { return dataClause.getAsyncOnlyAttr(); })
5497 .Default([&](mlir::Operation *) { return mlir::ArrayAttr{}; });
5498}
5499
5500std::optional<llvm::StringRef> mlir::acc::getVarName(mlir::Operation *accOp) {
5501 auto name{
5503 .Case<ACC_DATA_ENTRY_OPS, mlir::acc::MapInfoOp>(
5504 [&](auto entry) { return entry.getName(); })
5505 .Default([&](mlir::Operation *) -> std::optional<llvm::StringRef> {
5506 return {};
5507 })};
5508 return name;
5509}
5510
5511std::optional<mlir::acc::DataClause>
5513 auto dataClause{
5515 accDataEntryOp)
5516 .Case<ACC_DATA_ENTRY_OPS>(
5517 [&](auto entry) { return entry.getDataClause(); })
5518 .Default([&](mlir::Operation *) { return std::nullopt; })};
5519 return dataClause;
5520}
5521
5523 return llvm::TypeSwitch<mlir::Operation *, bool>(accDataEntryOp)
5524 .Case<ACC_DATA_ENTRY_OPS>([&](auto entry) { return entry.getImplicit(); })
5525 .Case<mlir::acc::MapInfoOp>([&](auto mapInfo) {
5526 return bitEnumContainsAny(mapInfo.getMapFlags(),
5527 mlir::acc::MapFlags::implicit);
5528 })
5529 .Default([&](mlir::Operation *) { return false; });
5530}
5531
5533 return llvm::TypeSwitch<mlir::Operation *, bool>(accDataClauseOp)
5534 .Case<ACC_DATA_CLAUSE_OPS>(
5535 [&](auto dataClause) { return dataClause.getSynthetic(); })
5536 .Default([&](mlir::Operation *) { return false; });
5537}
5538
5540 auto dataOperands{
5543 mlir::acc::KernelEnvironmentOp>(
5544 [&](auto entry) { return entry.getDataClauseOperands(); })
5545 .Default([&](mlir::Operation *) { return mlir::ValueRange(); })};
5546 return dataOperands;
5547}
5548
5551 auto dataOperands{
5554 mlir::acc::KernelEnvironmentOp>(
5555 [&](auto entry) { return entry.getDataClauseOperandsMutable(); })
5556 .Default([&](mlir::Operation *) { return nullptr; })};
5557 return dataOperands;
5558}
5559
5560mlir::SymbolRefAttr mlir::acc::getRecipe(mlir::Operation *accOp) {
5561 auto recipe{
5563 .Case<ACC_DATA_ENTRY_OPS>(
5564 [&](auto entry) { return entry.getRecipeAttr(); })
5565 .Default([&](mlir::Operation *) { return mlir::SymbolRefAttr{}; })};
5566 return recipe;
5567}
return success()
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static void printSourceLocation(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::LocationAttr locAttr)
Definition OpenACC.cpp:1042
void printRoutineGangClause(OpAsmPrinter &p, Operation *op, std::optional< mlir::ArrayAttr > gang, std::optional< mlir::ArrayAttr > gangDim, std::optional< mlir::ArrayAttr > gangDimDeviceTypes)
Definition OpenACC.cpp:4927
bool hasDuplicateDeviceTypes(std::optional< mlir::ArrayAttr > segments, llvm::SmallSet< mlir::acc::DeviceType, 3 > &deviceTypes)
Definition OpenACC.cpp:3647
static LogicalResult verifyDeviceTypeCountMatch(Op op, OperandRange operands, ArrayAttr deviceTypes, llvm::StringRef keyword)
Definition OpenACC.cpp:2219
static ParseResult parseArrayAttr(mlir::OpAsmParser &parser, mlir::ArrayAttr &attr)
Definition OpenACC.cpp:1069
static ParseResult parseBindName(OpAsmParser &parser, mlir::ArrayAttr &bindIdName, mlir::ArrayAttr &bindStrName, mlir::ArrayAttr &deviceIdTypes, mlir::ArrayAttr &deviceStrTypes)
Definition OpenACC.cpp:4778
static void printRecipeSym(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::SymbolRefAttr recipeAttr)
Definition OpenACC.cpp:1054
static mlir::Operation::operand_range getWaitValuesWithoutDevnum(std::optional< mlir::ArrayAttr > deviceTypeAttr, mlir::Operation::operand_range operands, std::optional< llvm::ArrayRef< int32_t > > segments, std::optional< mlir::ArrayAttr > hasWaitDevnum, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:807
static void printArrayAttr(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::ArrayAttr attr)
Definition OpenACC.cpp:1074
static bool hasOnlyDeviceTypeNone(std::optional< mlir::ArrayAttr > attrs)
Definition OpenACC.cpp:2745
static ParseResult parseRecipeSym(mlir::OpAsmParser &parser, mlir::SymbolRefAttr &recipeAttr)
Definition OpenACC.cpp:1047
static void printAccVar(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::Value accVar, mlir::Type accVarType)
Definition OpenACC.cpp:960
static mlir::Value getWaitDevnumValue(std::optional< mlir::ArrayAttr > deviceTypeAttr, mlir::Operation::operand_range operands, std::optional< llvm::ArrayRef< int32_t > > segments, std::optional< mlir::ArrayAttr > hasWaitDevnum, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:787
static bool hasAnyGangWorkerVectorForDeviceType(std::optional< mlir::ArrayAttr > numGangsDeviceType, mlir::Operation::operand_range numGangs, std::optional< llvm::ArrayRef< int32_t > > numGangsSegments, std::optional< mlir::ArrayAttr > numWorkersDeviceType, mlir::Operation::operand_range numWorkers, std::optional< mlir::ArrayAttr > vectorLengthDeviceType, mlir::Operation::operand_range vectorLength, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:2359
static void printVar(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::Value var)
Definition OpenACC.cpp:929
static void printWaitClause(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > deviceTypes, std::optional< mlir::DenseI32ArrayAttr > segments, std::optional< mlir::ArrayAttr > hasDevNum, std::optional< mlir::ArrayAttr > keywordOnly)
Definition OpenACC.cpp:2756
static ParseResult parseWaitClause(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::ArrayAttr &deviceTypes, mlir::DenseI32ArrayAttr &segments, mlir::ArrayAttr &hasDevNum, mlir::ArrayAttr &keywordOnly)
Definition OpenACC.cpp:2661
static BodyExecution getBodyExecution(LoopOp loopOp)
Prove whether the body of loopOp runs.
Definition OpenACC.cpp:632
static bool hasDeviceTypeValues(std::optional< mlir::ArrayAttr > arrayAttr)
Definition OpenACC.cpp:729
static void printDeviceTypeArrayAttr(mlir::OpAsmPrinter &p, mlir::Operation *op, std::optional< mlir::ArrayAttr > deviceTypes)
Definition OpenACC.cpp:4984
static ParseResult parseGangValue(OpAsmParser &parser, llvm::StringRef keyword, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, llvm::SmallVector< GangArgTypeAttr > &attributes, GangArgTypeAttr gangArgType, bool &needCommaBetweenValues, bool &newValue)
Definition OpenACC.cpp:3456
static ParseResult parseCombinedConstructsLoop(mlir::OpAsmParser &parser, mlir::acc::CombinedConstructsTypeAttr &attr)
Definition OpenACC.cpp:3002
static std::optional< mlir::acc::DeviceType > checkDeviceTypes(mlir::ArrayAttr deviceTypes)
Check for duplicates in the DeviceType array attribute.
Definition OpenACC.cpp:3663
static LogicalResult checkDeclareOperands(Op &op, const mlir::ValueRange &operands, bool requireAtLeastOneOperand=true)
Definition OpenACC.cpp:4682
static LogicalResult checkVarAndAccVar(Op op)
Definition OpenACC.cpp:867
static ParseResult parseOperandsWithKeywordOnly(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::UnitAttr &attr)
Definition OpenACC.cpp:2956
static void printDeviceTypes(mlir::OpAsmPrinter &p, std::optional< mlir::ArrayAttr > deviceTypes)
Definition OpenACC.cpp:747
static LogicalResult checkVarAndVarType(Op op)
Definition OpenACC.cpp:849
static LogicalResult checkValidModifier(Op op, acc::DataClauseModifier validModifiers)
Definition OpenACC.cpp:883
static void addOperandEffect(SmallVectorImpl< SideEffects::EffectInstance< MemoryEffects::Effect > > &effects, MutableOperandRange operand)
Helper to add an effect on an operand, referenced by its mutable range.
Definition OpenACC.cpp:1472
ParseResult parseLoopControl(OpAsmParser &parser, Region &region, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &lowerbound, SmallVectorImpl< Type > &lowerboundType, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &upperbound, SmallVectorImpl< Type > &upperboundType, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &step, SmallVectorImpl< Type > &stepType)
loop-control ::= control ( ssa-id-and-type-list ) = ( ssa-id-and-type-list ) to ( ssa-id-and-type-lis...
Definition OpenACC.cpp:4031
static LogicalResult checkDataOperands(Op op, const mlir::ValueRange &operands)
Check dataOperands for acc.parallel, acc.serial and acc.kernels.
Definition OpenACC.cpp:2174
static ParseResult parseDeviceTypeOperands(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::ArrayAttr &deviceTypes)
Definition OpenACC.cpp:2792
static mlir::Value getValueInDeviceTypeSegment(std::optional< mlir::ArrayAttr > arrayAttr, mlir::Operation::operand_range range, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:2302
static void addResultEffect(SmallVectorImpl< SideEffects::EffectInstance< MemoryEffects::Effect > > &effects, Value result)
Helper to add an effect on a result value.
Definition OpenACC.cpp:1482
static LogicalResult checkNoModifier(Op op)
Definition OpenACC.cpp:875
static ParseResult parseAccVar(mlir::OpAsmParser &parser, OpAsmParser::UnresolvedOperand &var, mlir::Type &accVarType)
Definition OpenACC.cpp:938
static std::optional< unsigned > findSegment(ArrayAttr segments, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:758
static ParseResult parseDenseBoolArrayAttr(mlir::OpAsmParser &parser, mlir::DenseBoolArrayAttr &attr)
Definition OpenACC.cpp:1059
static mlir::Operation::operand_range getValuesFromSegments(std::optional< mlir::ArrayAttr > arrayAttr, mlir::Operation::operand_range range, std::optional< llvm::ArrayRef< int32_t > > segments, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:771
static ParseResult parseNumGangs(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::ArrayAttr &deviceTypes, mlir::DenseI32ArrayAttr &segments)
Definition OpenACC.cpp:2531
static void getSingleRegionOpSuccessorRegions(Operation *op, Region &region, RegionBranchPoint point, SmallVectorImpl< RegionSuccessor > &regions)
Generic helper for single-region OpenACC ops that execute their body once and then continue after the...
Definition OpenACC.cpp:550
static ParseResult parseVar(mlir::OpAsmParser &parser, OpAsmParser::UnresolvedOperand &var)
Definition OpenACC.cpp:914
void printLoopControl(OpAsmPrinter &p, Operation *op, Region &region, ValueRange lowerbound, TypeRange lowerboundType, ValueRange upperbound, TypeRange upperboundType, ValueRange steps, TypeRange stepType)
Definition OpenACC.cpp:4062
static ValueRange getSingleRegionSuccessorInputs(Operation *op, RegionSuccessor successor)
Definition OpenACC.cpp:561
static void printDenseBoolArrayAttr(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::DenseBoolArrayAttr attr)
Definition OpenACC.cpp:1064
static ParseResult parseDeviceTypeArrayAttr(OpAsmParser &parser, mlir::ArrayAttr &deviceTypes)
Definition OpenACC.cpp:4957
static ParseResult parseRoutineGangClause(OpAsmParser &parser, mlir::ArrayAttr &gang, mlir::ArrayAttr &gangDim, mlir::ArrayAttr &gangDimDeviceTypes)
Definition OpenACC.cpp:4866
static void printDeviceTypeOperandsWithSegment(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > deviceTypes, std::optional< mlir::DenseI32ArrayAttr > segments)
Definition OpenACC.cpp:2644
static void printDeviceTypeOperands(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > deviceTypes)
Definition OpenACC.cpp:2819
static void printOperandWithKeywordOnly(mlir::OpAsmPrinter &p, mlir::Operation *op, std::optional< mlir::Value > operand, mlir::Type operandType, mlir::UnitAttr attr)
Definition OpenACC.cpp:2941
static ParseResult parseSourceLocation(mlir::OpAsmParser &parser, mlir::LocationAttr &locAttr)
Definition OpenACC.cpp:1030
static ParseResult parseDeviceTypeOperandsWithSegment(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::ArrayAttr &deviceTypes, mlir::DenseI32ArrayAttr &segments)
Definition OpenACC.cpp:2598
static bool isEnclosedIntoComputeOp(mlir::Operation *op)
Definition OpenACC.cpp:1466
static ParseResult parseOperandWithKeywordOnly(mlir::OpAsmParser &parser, std::optional< OpAsmParser::UnresolvedOperand > &operand, mlir::Type &operandType, mlir::UnitAttr &attr)
Definition OpenACC.cpp:2917
static void printVarPtrType(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::Type varPtrType, mlir::TypeAttr varTypeAttr)
Definition OpenACC.cpp:1004
static ParseResult parseGangClause(OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &gangOperands, llvm::SmallVectorImpl< Type > &gangOperandsType, mlir::ArrayAttr &gangArgType, mlir::ArrayAttr &deviceType, mlir::DenseI32ArrayAttr &segments, mlir::ArrayAttr &gangOnlyDeviceType)
Definition OpenACC.cpp:3475
static LogicalResult verifyInitLikeSingleArgRegion(Operation *op, Region &region, StringRef regionType, StringRef regionName, Type type, bool verifyYield, bool optional=false)
Definition OpenACC.cpp:1936
static void printOperandsWithKeywordOnly(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, mlir::UnitAttr attr)
Definition OpenACC.cpp:2986
static void printSingleDeviceType(mlir::OpAsmPrinter &p, mlir::Attribute attr)
Definition OpenACC.cpp:2575
static LogicalResult checkRecipe(OpT op, llvm::StringRef operandName)
Definition OpenACC.cpp:893
static LogicalResult checkPrivateOperands(mlir::Operation *accConstructOp, const mlir::ValueRange &operands, llvm::StringRef operandName)
Definition OpenACC.cpp:2188
static void printDeviceTypeOperandsWithKeywordOnly(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > deviceTypes, std::optional< mlir::ArrayAttr > keywordOnlyDeviceTypes)
Definition OpenACC.cpp:2898
static bool hasDeviceType(std::optional< mlir::ArrayAttr > arrayAttr, mlir::acc::DeviceType deviceType)
Definition OpenACC.cpp:733
void printGangClause(OpAsmPrinter &p, Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > gangArgTypes, std::optional< mlir::ArrayAttr > deviceTypes, std::optional< mlir::DenseI32ArrayAttr > segments, std::optional< mlir::ArrayAttr > gangOnlyDeviceTypes)
Definition OpenACC.cpp:3602
static ParseResult parseDeviceTypeOperandsWithKeywordOnly(mlir::OpAsmParser &parser, llvm::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &operands, llvm::SmallVectorImpl< Type > &types, mlir::ArrayAttr &deviceTypes, mlir::ArrayAttr &keywordOnlyDeviceType)
Definition OpenACC.cpp:2830
static ParseResult parseVarPtrType(mlir::OpAsmParser &parser, mlir::Type &varPtrType, mlir::TypeAttr &varTypeAttr)
Definition OpenACC.cpp:972
static LogicalResult checkWaitAndAsyncConflict(Op op)
Definition OpenACC.cpp:827
static LogicalResult verifyDeviceTypeAndSegmentCountMatch(Op op, OperandRange operands, DenseI32ArrayAttr segments, ArrayAttr deviceTypes, llvm::StringRef keyword, int32_t maxInSegment=0)
Definition OpenACC.cpp:2230
static unsigned getParallelismForDeviceType(acc::RoutineOp op, acc::DeviceType dtype)
Definition OpenACC.cpp:4744
static void printNumGangs(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRange operands, mlir::TypeRange types, std::optional< mlir::ArrayAttr > deviceTypes, std::optional< mlir::DenseI32ArrayAttr > segments)
Definition OpenACC.cpp:2581
BodyExecution
Whether the body of a structured acc.loop is proven to run.
Definition OpenACC.cpp:619
@ Always
The body runs at least once, so the parent cannot bypass the region.
Definition OpenACC.cpp:621
@ Never
The body never runs, so the parent cannot enter the region.
Definition OpenACC.cpp:623
@ Maybe
Neither could be proven, so the parent may do either.
Definition OpenACC.cpp:625
static void printCombinedConstructsLoop(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::acc::CombinedConstructsTypeAttr attr)
Definition OpenACC.cpp:3022
static void printBindName(mlir::OpAsmPrinter &p, mlir::Operation *op, std::optional< mlir::ArrayAttr > bindIdName, std::optional< mlir::ArrayAttr > bindStrName, std::optional< mlir::ArrayAttr > deviceIdTypes, std::optional< mlir::ArrayAttr > deviceStrTypes)
Definition OpenACC.cpp:4836
static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp)
ArrayAttr()
b getContext())
false
Parses a map_entries map type from a string format back into its numeric value.
static void replaceOpWithRegion(RewriterBase &rewriter, Operation *op, Region &region)
Replaces the given op with the contents of the given single-block region, using the operands of the b...
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 void genStore(OpBuilder &builder, Location loc, Value val, Value mem, Value idx)
Generates a store with proper index typing and proper value.
static Value genLoad(OpBuilder &builder, Location loc, Value mem, Value idx)
Generates a load with proper index typing.
virtual ParseResult parseLBrace()=0
Parse a { token.
@ 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 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 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 parseRSquare()=0
Parse a ] token.
virtual ParseResult parseRBrace()=0
Parse a } token.
virtual ParseResult parseOptionalRParen()=0
Parse a ) token if present.
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 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.
virtual ParseResult parseOptionalLSquare()=0
Parse a [ token if present.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual void printType(Type type)
virtual void printAttribute(Attribute attr)
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
iterator_range< args_iterator > addArguments(TypeRange types, ArrayRef< Location > locs)
Add one argument to the argument list for each type specified in the list.
Definition Block.cpp:165
Operation & front()
Definition Block.h:178
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgListType getArguments()
Definition Block.h:112
static BoolAttr get(MLIRContext *context, bool value)
IntegerType getI64Type()
Definition Builders.cpp:73
MLIRContext * getContext() const
Definition Builders.h:56
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
Location objects represent source locations information in MLIR.
Definition Location.h:32
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class provides a mutable adaptor for a range of operands.
Definition ValueRange.h:119
unsigned size() const
Returns the current size of the range.
Definition ValueRange.h:157
void append(ValueRange values)
Append the given values to the range.
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 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 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 printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
virtual void printOperand(Value value)=0
Print implementations for various things an operation contains.
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Definition Builders.cpp:439
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Location getLoc()
The source location the operation was defined or derived from.
This provides public APIs that all operations should have.
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
Definition Operation.h:512
OperandRange operand_range
Definition Operation.h:396
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
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 represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
This class represents a successor of a region.
bool isOperation() const
Return true if the successor is an operation.
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
iterator_range< OpIterator > getOps()
Definition Region.h:180
bool empty()
Definition Region.h:60
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
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 specific instance of an effect.
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
bool isIntOrIndexOrFloat() const
Return true if this is an integer (of any signedness), index, or float type.
Definition Types.cpp:122
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
Base attribute class for language-specific variable information carried through the OpenACC type inte...
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
#define ACC_COMPUTE_CONSTRUCT_OPS
Definition OpenACC.h:63
#define ACC_COMPUTE_AND_DATA_CONSTRUCT_OPS
Definition OpenACC.h:74
#define ACC_DATA_CLAUSE_OPS
Definition OpenACC.h:62
#define ACC_DATA_ENTRY_OPS
Definition OpenACC.h:49
#define ACC_DATA_EXIT_OPS
Definition OpenACC.h:59
bool getSyntheticFlag(mlir::Operation *accDataClauseOp)
Used to find out whether the implementation created the data operation for its own bookkeeping,...
Definition OpenACC.cpp:5532
mlir::Value getAccVar(mlir::Operation *accDataClauseOp)
Used to obtain the accVar from a data clause operation.
Definition OpenACC.cpp:5435
mlir::Value getVar(mlir::Operation *accDataClauseOp)
Used to obtain the var from a data clause operation.
Definition OpenACC.cpp:5403
mlir::TypedValue< mlir::acc::PointerLikeType > getAccPtr(mlir::Operation *accDataClauseOp)
Used to obtain the accVar from a data clause operation if it implements PointerLikeType.
Definition OpenACC.cpp:5422
std::optional< mlir::acc::DataClause > getDataClause(mlir::Operation *accDataEntryOp)
Used to obtain the dataClause from a data entry operation.
Definition OpenACC.cpp:5512
mlir::MutableOperandRange getMutableDataOperands(mlir::Operation *accOp)
Used to get a mutable range iterating over the data operands.
Definition OpenACC.cpp:5550
mlir::SmallVector< mlir::Value > getBounds(mlir::Operation *accDataClauseOp)
Used to obtain bounds from an acc data clause operation.
Definition OpenACC.cpp:5454
std::optional< ClauseDefaultValue > getDefaultAttr(mlir::Operation *op)
Looks for an OpenACC default attribute on the current operation op or in a parent operation which enc...
bool hasWaitDevnum(OpTy op, DeviceType deviceType)
Returns whether the wait clause op gives for deviceType carries a devnum modifier,...
mlir::ValueRange getDataOperands(mlir::Operation *accOp)
Used to get an immutable range iterating over the data operands.
Definition OpenACC.cpp:5539
std::optional< llvm::StringRef > getVarName(mlir::Operation *accOp)
Used to obtain the name from an acc operation.
Definition OpenACC.cpp:5500
bool isGangWorkerVectorAllOne(ComputeOpT op)
Definition OpenACC.h:251
bool getImplicitFlag(mlir::Operation *accDataEntryOp)
Used to find out whether data operation is implicit.
Definition OpenACC.cpp:5522
mlir::SymbolRefAttr getRecipe(mlir::Operation *accOp)
Used to get the recipe attribute from a data clause operation.
Definition OpenACC.cpp:5560
mlir::SmallVector< mlir::Value > getAsyncOperands(mlir::Operation *accDataClauseOp)
Used to obtain async operands from an acc data clause operation.
Definition OpenACC.cpp:5471
bool isMappableType(mlir::Type type)
Used to check whether the provided type implements the MappableType interface.
Definition OpenACC.h:180
mlir::Value getVarPtrPtr(mlir::Operation *accDataClauseOp)
Used to obtain the varPtrPtr from a data clause operation.
Definition OpenACC.cpp:5444
static constexpr StringLiteral getVarNameAttrName()
Definition OpenACC.h:224
mlir::ArrayAttr getAsyncOnly(mlir::Operation *accDataClauseOp)
Returns an array of acc:DeviceTypeAttr attributes attached to an acc data clause operation,...
Definition OpenACC.cpp:5493
mlir::Type getVarType(mlir::Operation *accDataClauseOp)
Used to obtains the varType from a data clause operation which records the type of variable.
Definition OpenACC.cpp:5411
mlir::TypedValue< mlir::acc::PointerLikeType > getVarPtr(mlir::Operation *accDataClauseOp)
Used to obtain the var from a data clause operation if it implements PointerLikeType.
Definition OpenACC.cpp:5389
mlir::ArrayAttr getAsyncOperandsDeviceType(mlir::Operation *accDataClauseOp)
Returns an array of acc:DeviceTypeAttr attributes attached to an acc data clause operation,...
Definition OpenACC.cpp:5485
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Value genCast(OpBuilder &builder, Location loc, Value value, Type dstTy)
Add type casting between arith and index types when needed.
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
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
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
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
detail::DenseArrayAttrImpl< bool > DenseBoolArrayAttr
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
This is the representation of an operand reference.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
This represents an operation in an abstracted form, suitable for use with the builder APIs.