26#include "llvm/ADT/TypeSwitch.h"
27#include "llvm/Support/Debug.h"
34#define DEBUG_TYPE "acc-atomic-patterns"
43template <
typename AtomicOpTy>
52 accSupport(accSupport), getLoadAddress(getLoadAddress) {}
55 matchAndRewrite(AtomicOpTy op, OpAdaptor adaptor,
56 ConversionPatternRewriter &rewriter)
const override;
62 size_t getComplexStructElementSizeInBits(
Type ty)
const {
63 auto structType = dyn_cast<LLVM::LLVMStructType>(ty);
64 if (!structType || structType.getBody().size() != 2 ||
65 structType.getBody()[0] != structType.getBody()[1] ||
66 !structType.getBody()[0].isIntOrFloat())
68 size_t elementSizeInBits = structType.getBody()[0].getIntOrFloatBitWidth();
69 if (elementSizeInBits == 0 || elementSizeInBits > 64)
70 llvm_unreachable(
"unexpected complex type");
71 return elementSizeInBits;
75 Value serializeExpr(
Value expr, ConversionPatternRewriter &rewriter)
const {
76 size_t elementSizeInBits =
77 getComplexStructElementSizeInBits(expr.
getType());
78 if (!elementSizeInBits)
82 Value firstValue = LLVM::ExtractValueOp::create(rewriter, loc, expr, 0);
83 Value secondValue = LLVM::ExtractValueOp::create(rewriter, loc, expr, 1);
85 Type intCastTy = IntegerType::get(context, elementSizeInBits);
86 Type intTy = IntegerType::get(context, elementSizeInBits * 2);
87 firstValue = LLVM::BitcastOp::create(rewriter, loc, intCastTy, firstValue);
89 LLVM::BitcastOp::create(rewriter, loc, intCastTy, secondValue);
90 firstValue = LLVM::ZExtOp::create(rewriter, loc, intTy, firstValue);
91 secondValue = LLVM::ZExtOp::create(rewriter, loc, intTy, secondValue);
93 LLVM::ConstantOp::create(rewriter, loc, intTy, elementSizeInBits);
94 Value result = LLVM::ShlOp::create(rewriter, loc, secondValue, shlVal);
95 return LLVM::OrOp::create(rewriter, loc,
result, firstValue);
100 ConversionPatternRewriter &rewriter)
const {
101 size_t elementSizeInBits = getComplexStructElementSizeInBits(origTy);
102 if (!elementSizeInBits)
105 auto intTy = dyn_cast<IntegerType>(expr.
getType());
106 if (!intTy || intTy.getWidth() != elementSizeInBits * 2)
111 Type elemTy = IntegerType::get(context, elementSizeInBits);
112 Value low = LLVM::TruncOp::create(rewriter, loc, elemTy, expr);
114 LLVM::ConstantOp::create(rewriter, loc, intTy, elementSizeInBits);
115 Value highFull = LLVM::LShrOp::create(rewriter, loc, expr, shiftAmount);
116 Value high = LLVM::TruncOp::create(rewriter, loc, elemTy, highFull);
117 Type origElemTy = cast<LLVM::LLVMStructType>(origTy).getBody()[0];
118 if (origElemTy != elemTy) {
119 low = LLVM::BitcastOp::create(rewriter, loc, origElemTy, low);
120 high = LLVM::BitcastOp::create(rewriter, loc, origElemTy, high);
122 Value undefStruct = LLVM::UndefOp::create(rewriter, loc, origTy);
123 Value structWithLow = LLVM::InsertValueOp::create(
125 return LLVM::InsertValueOp::create(rewriter, loc, origTy, structWithLow,
129 uint64_t getAtomicSizeInBytes(
Type originalTy,
Type convertedTy,
130 ModuleOp module)
const {
131 std::optional<TypeSizeAndAlignment> sizeAndAlignment =
133 if (!sizeAndAlignment)
136 assert(sizeAndAlignment &&
"atomic type size is not computable");
137 return sizeAndAlignment->first.getFixedValue();
140 Type getAtomicType(
Type originalTy,
Type convertedTy, ModuleOp module)
const {
143 return IntegerType::get(
145 getAtomicSizeInBytes(originalTy, convertedTy, module) *
kBitsInByte);
148 Type getReferencedElementType(
Value ref, ModuleOp module =
nullptr)
const {
149 auto ptr = dyn_cast<PointerLikeType>(ref.
getType());
151 llvm_unreachable(
"unexpected type");
152 Type elementTy =
ptr.getElementType();
153 Type convertedTy = this->getTypeConverter()->convertType(elementTy);
155 return getAtomicType(elementTy, convertedTy, module);
162 ConversionPatternRewriter &rewriter)
const {
163 auto memrefTy = dyn_cast<MemRefType>(originalRef.
getType());
169 LLVM::ExtractValueOp::create(rewriter, loc, convertedPtr, 1);
170 Value offset = LLVM::ExtractValueOp::create(rewriter, loc, convertedPtr, 2);
171 Type elemPtrType = LLVM::LLVMPointerType::get(rewriter.getContext());
172 return LLVM::GEPOp::create(
173 rewriter, loc, elemPtrType,
174 this->getTypeConverter()->convertType(memrefTy.getElementType()),
179 ConversionPatternRewriter &rewriter)
const;
184 tryEmitCaptureAtomicRMW(AtomicCaptureOp capture, AtomicUpdateOp update,
186 ConversionPatternRewriter &rewriter)
const;
188 Value genUpdateCmpxchgLoop(AtomicUpdateOp update,
189 ConversionPatternRewriter &rewriter)
const;
193LogicalResult ACCAtomicOpConversion<AtomicReadOp>::matchAndRewrite(
194 AtomicReadOp read, OpAdaptor adaptor,
195 ConversionPatternRewriter &rewriter)
const {
197 Value xRef = read.getX();
198 Value xPtr = getAtomicPointer(xRef, adaptor.getX(), loc, rewriter);
199 ModuleOp mod = read->getParentOfType<ModuleOp>();
200 Type xType = getReferencedElementType(xRef, mod);
202 auto ordering = LLVM::AtomicOrdering::monotonic;
204 if (xType.isSignlessInteger()) {
205 Value zero = LLVM::ConstantOp::create(rewriter, loc, xType, 0);
206 storeVal = LLVM::AtomicRMWOp::create(rewriter, loc, LLVM::AtomicBinOp::_or,
207 xPtr, zero, ordering);
209 unsigned bitWidth = xType.getIntOrFloatBitWidth();
210 Type intType = IntegerType::get(rewriter.getContext(), bitWidth);
211 Value zero = LLVM::ConstantOp::create(rewriter, loc, intType, 0);
212 Value intVal = LLVM::AtomicRMWOp::create(
213 rewriter, loc, LLVM::AtomicBinOp::_or, xPtr, zero, ordering);
214 storeVal = LLVM::BitcastOp::create(rewriter, loc, xType, intVal);
217 Value vRef = read.getV();
218 Value vPtr = getAtomicPointer(vRef, adaptor.getV(), loc, rewriter);
219 Type vType = getReferencedElementType(vRef, mod);
221 if (xType != vType) {
223 auto vPtrType = cast<PointerLikeType>(vRef.
getType());
224 storeVal = vPtrType.genCast(rewriter, loc, storeVal, vType);
226 return rewriter.notifyMatchFailure(
227 read,
"failed to convert the loaded value to the destination type");
230 auto storeOp = LLVM::StoreOp::create(rewriter, loc, storeVal, vPtr);
231 rewriter.replaceOp(read, storeOp);
236LogicalResult ACCAtomicOpConversion<AtomicWriteOp>::matchAndRewrite(
237 AtomicWriteOp write, OpAdaptor adaptor,
238 ConversionPatternRewriter &rewriter)
const {
240 Value expr = serializeExpr(adaptor.getExpr(), rewriter);
241 Value xRef = write.getX();
242 Value xPtr = getAtomicPointer(xRef, adaptor.getX(), loc, rewriter);
243 ModuleOp mod = write->getParentOfType<ModuleOp>();
244 Type xType = getReferencedElementType(xRef, mod);
246 auto ordering = LLVM::AtomicOrdering::monotonic;
247 if (!xType.isSignlessInteger()) {
248 unsigned bitWidth = xType.getIntOrFloatBitWidth();
249 Type intType = IntegerType::get(rewriter.getContext(), bitWidth);
250 expr = LLVM::BitcastOp::create(rewriter, loc, intType, expr);
252 LLVM::AtomicRMWOp::create(rewriter, loc, LLVM::AtomicBinOp::xchg, xPtr, expr,
254 rewriter.eraseOp(write);
259template <
typename AtomicOpTy>
260Block *ACCAtomicOpConversion<AtomicOpTy>::constructCmpxchgLoop(
262 ConversionPatternRewriter &rewriter)
const {
264 Location loc = rewriter.getInsertionPoint()->getLoc();
265 Block *initBlock = rewriter.getInsertionBlock();
267 rewriter.splitBlock(initBlock, rewriter.getInsertionPoint());
270 rewriter.splitBlock(loopBlock, rewriter.getInsertionPoint());
273 rewriter.setInsertionPointToEnd(initBlock);
274 Value init = LLVM::LoadOp::create(rewriter, loc, type,
ptr);
275 LLVM::BrOp::create(rewriter, loc, init, loopBlock);
278 rewriter.setInsertionPointToStart(loopBlock);
281 Value result = serializeExpr(rewriter.getRemappedValue(expr), rewriter);
283 Type convertedExprType =
284 this->getTypeConverter()->convertType(expr.
getType());
285 Type exprType = getAtomicType(expr.
getType(), convertedExprType, mod);
289 Type tmpType = IntegerType::get(
290 rewriter.getContext(),
291 getAtomicSizeInBytes(expr.
getType(), convertedExprType, mod) *
293 result = LLVM::BitcastOp::create(rewriter, loc, tmpType,
result);
295 LLVM::BitcastOp::create(rewriter, loc, tmpType, loopArgument);
300 auto successOrdering = LLVM::AtomicOrdering::acq_rel;
301 auto failureOrdering = LLVM::AtomicOrdering::monotonic;
303 LLVM::AtomicCmpXchgOp::create(rewriter, loc,
ptr, loopArgument,
result,
304 successOrdering, failureOrdering);
306 Value newLoaded = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 0);
307 Value ok = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 1);
311 newLoaded = LLVM::BitcastOp::create(rewriter, loc, exprType, newLoaded);
315 loopBlock, newLoaded);
320static Value skipUnrealizedConversionOp(
Value v) {
322 dyn_cast_or_null<UnrealizedConversionCastOp>(v.
getDefiningOp()))
323 return skipUnrealizedConversionOp(convOp.getOperand(0));
328static Value getBaseStorage(
Value addr, ConversionPatternRewriter &rewriter) {
329 Operation *op = skipUnrealizedConversionOp(addr).getDefiningOp();
330 if (
auto gepOp = dyn_cast_or_null<LLVM::GEPOp>(op)) {
331 addr = gepOp.getBase();
334 if (
auto extractOp = dyn_cast_or_null<LLVM::ExtractValueOp>(op)) {
337 ArrayRef extractingPosition = extractOp.getPosition();
338 Value container = skipUnrealizedConversionOp(extractOp.getContainer());
339 Value remappedContainer = rewriter.getRemappedValue(container);
340 op = remappedContainer ? remappedContainer.
getDefiningOp() :
nullptr;
342 while (
auto insertValueOp = dyn_cast_or_null<LLVM::InsertValueOp>(op)) {
343 ArrayRef insertingPosition = insertValueOp.getPosition();
344 if (insertingPosition == extractingPosition) {
345 addr = insertValueOp.getValue();
349 op = insertValueOp.getContainer().getDefiningOp();
357 return skipUnrealizedConversionOp(addr);
364 ConversionPatternRewriter &rewriter,
366 Value vStorage = getBaseStorage(vPtr, rewriter);
369 Value vStorageRef = skipUnrealizedConversionOp(vRef);
371 llvm::dbgs() <<
"[acc-atomic] moveDependency\n";
372 llvm::dbgs() <<
" vPtr = " << vPtr <<
"\n";
373 llvm::dbgs() <<
" vStorage = " << vStorage <<
"\n";
374 llvm::dbgs() <<
" vStorageRef = " << vStorageRef <<
"\n";
375 llvm::dbgs() <<
" expr = " << expr <<
"\n";
378 std::set<std::pair<Operation *, std::set<Operation *>>> worklist;
379 std::set<Operation *> included;
380 Value mappedExpr = rewriter.getRemappedValue(expr);
381 remappedToOriginal[mappedExpr] = expr;
383 worklist.insert(std::pair{exprDef, std::set<Operation *>{}});
385 while (!worklist.empty()) {
386 auto [dep, post] = worklist.extract(worklist.begin()).value();
389 LLVM_DEBUG(llvm::dbgs() <<
" visit dep = " << *dep <<
"\n");
393 LLVM_DEBUG(llvm::dbgs() <<
" -> skipped: not in parental region\n");
398 if (
auto load = dyn_cast<LLVM::LoadOp>(dep)) {
399 addr =
load.getAddr();
400 }
else if (
auto load = dyn_cast<memref::LoadOp>(dep)) {
401 addr =
load.getMemref();
402 }
else if (getLoadAddress) {
403 addr = getLoadAddress(dep);
406 Value baseStorage = getBaseStorage(addr, rewriter);
408 llvm::dbgs() <<
" addr = " << addr <<
"\n";
409 llvm::dbgs() <<
" baseStorage = " << baseStorage <<
"\n";
410 llvm::dbgs() <<
" remapped = "
411 << rewriter.getRemappedValue(baseStorage) <<
"\n";
412 llvm::dbgs() <<
" matchRef=" << (baseStorage == vStorageRef)
414 << (rewriter.getRemappedValue(baseStorage) == vStorage)
417 if (baseStorage == vStorageRef ||
418 rewriter.getRemappedValue(baseStorage) == vStorage) {
422 rewriter.modifyOpInPlace(op, [&] {
423 op->replaceUsesOfWith(remappedToOriginal[
load], loopArgument);
428 replaceUses(&loopHead);
429 included.insert(post.begin(), post.end());
434 for (
Value operand : dep->getOperands()) {
435 Value mappedOperand = rewriter.getRemappedValue(operand);
438 if (
auto *op = operand.getDefiningOp()) {
439 remappedToOriginal[op->getResult(0)] = operand;
440 worklist.insert(std::pair{op, post});
443 remappedToOriginal[mappedOperand] = operand;
444 worklist.insert(std::pair{d, post});
453 if (included.find(op) != included.end())
454 includedInOrder.push_back(op);
457 d->moveBefore(&loopHead);
461template <
typename AtomicOpTy>
462Value ACCAtomicOpConversion<AtomicOpTy>::genUpdateCmpxchgLoop(
463 AtomicUpdateOp update, ConversionPatternRewriter &rewriter)
const {
465 Value xRef = update.getX();
467 getAtomicPointer(xRef, rewriter.getRemappedValue(xRef), loc, rewriter);
468 ModuleOp mod = update->getParentOfType<ModuleOp>();
469 Type xType = getReferencedElementType(xRef, mod);
470 Type xTypeOrig = getReferencedElementType(xRef);
472 Block &updateBlock = update.getRegion().
front();
477 Block *loopBlock = constructCmpxchgLoop(xPtr, xType, expr, rewriter);
481 rewriter.setInsertionPointToStart(loopBlock);
482 loopArgument = deserializeExpr(loopArgument, xTypeOrig, rewriter);
486 moveDependency(xRef, xPtr, expr, loopArgument, loopHead, rewriter,
489 rewriter.replaceAllUsesWith(cast<BlockArgument>(updateArgument),
494 rewriter.moveOpBefore(op, &loopHead);
497 if (
auto cmpxchg = dyn_cast<LLVM::AtomicCmpXchgOp>(loopHead))
498 return cmpxchg.getVal();
499 if (
auto bitcast = dyn_cast<LLVM::BitcastOp>(loopHead))
501 return bitcast.getArg();
502 if (
auto extract = dyn_cast<LLVM::ExtractValueOp>(loopHead))
504 return extract.getContainer();
505 llvm_unreachable(
"invalid cmpxchg loop");
508static std::optional<LLVM::AtomicBinOp> getAtomicBinOp(
Operation *op,
511 .Case<arith::AddFOp>([](
auto) {
return LLVM::AtomicBinOp::fadd; })
512 .Case<arith::AddIOp>([](
auto) {
return LLVM::AtomicBinOp::add; })
513 .Case<arith::SubFOp>(
514 [updateIsLhs](
auto) -> std::optional<LLVM::AtomicBinOp> {
518 return LLVM::AtomicBinOp::fsub;
520 .Case<arith::SubIOp>(
521 [updateIsLhs](
auto) -> std::optional<LLVM::AtomicBinOp> {
525 return LLVM::AtomicBinOp::sub;
527 .Case<arith::AndIOp>([](
auto) {
return LLVM::AtomicBinOp::_and; })
528 .Case<arith::OrIOp>([](
auto) {
return LLVM::AtomicBinOp::_or; })
529 .Case<arith::XOrIOp>([](
auto) {
return LLVM::AtomicBinOp::_xor; })
530 .Case<arith::MaxSIOp>([](
auto) {
return LLVM::AtomicBinOp::max; })
531 .Case<arith::MinSIOp>([](
auto) {
return LLVM::AtomicBinOp::min; })
532 .Case<arith::MaxUIOp>([](
auto) {
return LLVM::AtomicBinOp::umax; })
533 .Case<arith::MinUIOp>([](
auto) {
return LLVM::AtomicBinOp::umin; })
534 .Case<arith::MaximumFOp>([](
auto) {
return LLVM::AtomicBinOp::fmaximum; })
535 .Case<arith::MinimumFOp>([](
auto) {
return LLVM::AtomicBinOp::fminimum; })
536 .Case<arith::MaxNumFOp>(
537 [](
auto) {
return LLVM::AtomicBinOp::fmaximumnum; })
538 .Case<arith::MinNumFOp>(
539 [](
auto) {
return LLVM::AtomicBinOp::fminimumnum; })
540 .Default([](
Operation *) {
return std::nullopt; });
547static std::optional<std::pair<LLVM::AtomicBinOp, Operation *>>
548matchAtomicBinOpUpdate(AtomicUpdateOp update) {
561 bool updateIsLhs = binOp->
getOperand(0) == arg;
562 std::optional<LLVM::AtomicBinOp> kind = getAtomicBinOp(binOp, updateIsLhs);
565 return std::make_pair(*kind, binOp);
570LogicalResult ACCAtomicOpConversion<AtomicUpdateOp>::matchAndRewrite(
571 AtomicUpdateOp update, OpAdaptor adaptor,
572 ConversionPatternRewriter &rewriter)
const {
573 Block &updateBlock = update.getRegion().
front();
577 std::set<Operation *> dependents;
579 worklist.push_back(updateArgument);
580 while (!worklist.empty()) {
581 Value value = worklist.back();
585 dependents.insert(useOp);
594 if (dependents.find(&op) == dependents.end()) {
596 llvm_unreachable(
"invalid update operation");
597 independent.push_back(&op);
601 rewriter.moveOpBefore(op, update);
622 std::optional<Value> val = std::nullopt;
623 std::optional<LLVM::AtomicBinOp> kind = std::nullopt;
625 if (
auto matched = matchAtomicBinOpUpdate(update)) {
627 bool updateIsLhs = binOp->
getOperand(0) == updateArgument;
628 kind = matched->first;
635 struct ComponentAtomic {
636 LLVM::AtomicBinOp kind;
643 Type convertedArgTy =
644 this->getTypeConverter()->convertType(updateArgument.
getType());
645 if (
auto structTy = dyn_cast<LLVM::LLVMStructType>(convertedArgTy)) {
646 if (structTy.getBody().size() == 2 &&
647 structTy.getBody()[0] == structTy.getBody()[1] &&
648 structTy.getBody()[0].isIntOrFloat() &&
649 structTy.getBody()[0].getIntOrFloatBitWidth() > 32) {
653 int32_t fieldIdx = -1;
654 Value externalVal =
nullptr;
655 bool updateIsLhs =
false;
656 for (
unsigned i = 0; i < 2; ++i) {
659 if (reOp.getOperand() == updateArgument) {
662 updateIsLhs = i == 0;
664 }
else if (
auto imOp = operand.
getDefiningOp<complex::ImOp>()) {
665 if (imOp.getOperand() == updateArgument) {
668 updateIsLhs = i == 0;
674 auto componentKind = getAtomicBinOp(&op, updateIsLhs);
677 componentAtomics.push_back({*componentKind, externalVal, fieldIdx});
684 Value xPtr = getAtomicPointer(update.getX(), adaptor.getX(), loc, rewriter);
688 bool hasDistinctComplexLanes =
false;
689 if (componentAtomics.size() == 2) {
691 for (
const ComponentAtomic &ca : componentAtomics) {
692 if (ca.fieldIdx == 0 || ca.fieldIdx == 1)
693 lanes |= 1u << ca.fieldIdx;
695 hasDistinctComplexLanes = lanes == 0b11;
699 auto ordering = LLVM::AtomicOrdering::monotonic;
700 LLVM::AtomicRMWOp::create(rewriter, loc, *kind, xPtr,
701 rewriter.getRemappedValue(*val), ordering);
702 }
else if (hasDistinctComplexLanes) {
703 auto structTy = cast<LLVM::LLVMStructType>(
704 this->getTypeConverter()->convertType(updateArgument.
getType()));
705 Type ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
706 auto ordering = LLVM::AtomicOrdering::monotonic;
707 for (ComponentAtomic &ca : componentAtomics) {
709 LLVM::GEPOp::create(rewriter, loc, ptrType, structTy, xPtr,
711 LLVM::AtomicRMWOp::create(rewriter, loc, ca.kind, elemPtr,
712 rewriter.getRemappedValue(ca.val), ordering);
716 genUpdateCmpxchgLoop(update, rewriter);
718 rewriter.eraseOp(update);
726static bool exprReadsMemory(
Value expr) {
729 while (!worklist.empty()) {
730 Operation *def = worklist.pop_back_val().getDefiningOp();
731 if (!def || !seen.insert(def).second)
741template <
typename AtomicOpTy>
742LogicalResult ACCAtomicOpConversion<AtomicOpTy>::tryEmitCaptureAtomicRMW(
743 AtomicCaptureOp capture, AtomicUpdateOp update, AtomicReadOp read,
744 ConversionPatternRewriter &rewriter)
const {
745 if (read.getX() != update.getX())
747 auto matched = matchAtomicBinOpUpdate(update);
750 auto [kind, binOp] = *matched;
752 Value arg = update.getRegion().front().getArgument(0);
753 bool updateIsLhs = binOp->getOperand(0) == arg;
754 Value expr = binOp->getOperand(updateIsLhs ? 1 : 0);
759 this->getTypeConverter()->convertType(argTy) != argTy)
763 if (exprDef && exprDef->
getBlock() == &update.getRegion().
front())
765 if (exprReadsMemory(expr))
769 Value xRef = update.getX();
770 Value vRef = read.getV();
772 getAtomicPointer(xRef, rewriter.getRemappedValue(xRef), loc, rewriter);
774 getAtomicPointer(vRef, rewriter.getRemappedValue(vRef), loc, rewriter);
776 rewriter.setInsertionPoint(capture);
777 auto rmw = LLVM::AtomicRMWOp::create(rewriter, loc, kind, xPtr,
778 rewriter.getRemappedValue(expr),
779 LLVM::AtomicOrdering::monotonic);
781 Value captured = rmw.getRes();
782 if (capture.getFirstOp() == update.getOperation()) {
783 rewriter.moveOpAfter(binOp, rmw);
784 binOp->replaceUsesOfWith(arg, captured);
785 captured = binOp->getResult(0);
787 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(capture, captured, vPtr);
793LogicalResult ACCAtomicOpConversion<AtomicCaptureOp>::matchAndRewrite(
794 AtomicCaptureOp capture, OpAdaptor ,
795 ConversionPatternRewriter &rewriter)
const {
796 Operation *firstOp = capture.getFirstOp();
797 Operation *secondOp = capture.getSecondOp();
798 Value vPtr =
nullptr;
799 Value storeVal =
nullptr;
803 if (AtomicUpdateOp update = capture.getAtomicUpdateOp())
804 if (AtomicReadOp read = capture.getAtomicReadOp())
805 if (succeeded(tryEmitCaptureAtomicRMW(capture, update, read, rewriter)))
808 if (
auto firstReadStmt = dyn_cast<AtomicReadOp>(firstOp)) {
810 Value xRef = firstReadStmt.getX();
812 getAtomicPointer(xRef, rewriter.getRemappedValue(xRef), loc, rewriter);
813 ModuleOp mod = capture->getParentOfType<ModuleOp>();
814 Type xType = getReferencedElementType(xRef, mod);
815 Type xTypeOrig = getReferencedElementType(xRef);
816 Value vRef = firstReadStmt.getV();
818 getAtomicPointer(vRef, rewriter.getRemappedValue(vRef), loc, rewriter);
820 Value expr =
nullptr;
821 if (
auto secondWriteStmt = dyn_cast<AtomicWriteOp>(secondOp)) {
823 expr = secondWriteStmt.getExpr();
824 }
else if (
auto secondUpdateStmt = dyn_cast<AtomicUpdateOp>(secondOp)) {
826 Block &updateBlock = secondUpdateStmt.getRegion().
front();
831 Block *loopBlock = constructCmpxchgLoop(xPtr, xType, expr, rewriter);
834 auto condBr = cast<LLVM::CondBrOp>(loopBlock->
back());
835 storeVal = condBr.getFalseDestOperands()[0];
837 rewriter.setInsertionPointToStart(loopBlock);
838 loopArgument = deserializeExpr(loopArgument, xTypeOrig, rewriter);
841 if (
auto secondUpdateStmt = dyn_cast<AtomicUpdateOp>(secondOp)) {
842 Block &updateBlock = secondUpdateStmt.getRegion().
front();
847 rewriter.moveOpBefore(op, &loopHead);
850 moveDependency(vRef, vPtr, expr, loopArgument, loopHead, rewriter,
852 }
else if (
auto firstUpdateStmt = dyn_cast<AtomicUpdateOp>(firstOp)) {
853 if (
auto secondReadStmt = dyn_cast<AtomicReadOp>(secondOp)) {
855 storeVal = genUpdateCmpxchgLoop(firstUpdateStmt, rewriter);
857 Value vRef = secondReadStmt.getV();
858 vPtr = getAtomicPointer(vRef, rewriter.getRemappedValue(vRef),
859 capture.getLoc(), rewriter);
863 rewriter.setInsertionPoint(capture);
864 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(capture, storeVal, vPtr);
873 target.addIllegalOp<AtomicReadOp, AtomicWriteOp, AtomicUpdateOp,
881 patterns.
add<ACCAtomicOpConversion<AtomicReadOp>,
882 ACCAtomicOpConversion<AtomicWriteOp>,
883 ACCAtomicOpConversion<AtomicUpdateOp>,
884 ACCAtomicOpConversion<AtomicCaptureOp>>(converter, accSupport,
constexpr static const uint64_t kBitsInByte
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
OpListType & getOperations()
RetT walk(FnT &&callback)
Walk all nested operations, blocks (including this block) or regions, depending on the type of callba...
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
typename SourceOp::Adaptor OpAdaptor
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class represents an operand of an operation.
This class provides the API for ops that are known to be terminators.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Block * getBlock()
Returns the operation block that contains this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
unsigned getNumOperands()
operand_range getOperands()
Returns an iterator on the underlying Value's.
bool isAncestor(Operation *other)
Return true if this operation is an ancestor of the other operation.
unsigned getNumResults()
Return the number of results held by this operation.
ParentT getParentOfType()
Find the first parent operation of the given type, or nullptr if there is no ancestor operation.
RetT walk(FnT &&callback)
Walk all nested operations, blocks or regions (including this region), depending on the type of callb...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
void replaceAllUsesWith(Value newValue)
Replace all uses of 'this' value with the new value, updating anything in the IR that uses 'this' to ...
bool hasOneUse() const
Returns true if this value has exactly one use.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Region * getParentRegion()
Return the Region in which this Value is defined.
use_iterator use_begin() const
std::optional< TypeSizeAndAlignment > getTypeSizeAndAlignment(Type ty, ModuleOp module, const DataLayout &dl, OpenACCSupport *support=nullptr, Value var={})
Returns the size and ABI alignment in bytes.
Include the generated interface declarations.
void populateACCAtomicPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, acc::OpenACCSupport &accSupport, ACCAtomicLoadAddressCallback getLoadAddress={})
Populate patterns that lower OpenACC atomic operations to LLVM dialect.
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
llvm::TypeSwitch< T, ResultT > TypeSwitch
std::function< Value(Operation *)> ACCAtomicLoadAddressCallback
Returns the address operand of a dialect-specific load operation.
void configureACCAtomicConversionLegality(ConversionTarget &target)
Configure conversion legality for OpenACC atomic operations.