1629 lines · cpp
1//===- LLVMMemorySlot.cpp - MemorySlot interfaces ---------------*- C++ -*-===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file implements MemorySlot-related interfaces for LLVM dialect10// operations.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Dialect/LLVMIR/LLVMDialect.h"15#include "mlir/Dialect/LLVMIR/LLVMTypes.h"16#include "mlir/IR/Matchers.h"17#include "mlir/IR/PatternMatch.h"18#include "mlir/Interfaces/DataLayoutInterfaces.h"19#include "mlir/Interfaces/MemorySlotInterfaces.h"20#include "llvm/ADT/STLExtras.h"21#include "llvm/ADT/TypeSwitch.h"22#include "llvm/Support/DebugLog.h"23 24#define DEBUG_TYPE "sroa"25 26using namespace mlir;27 28//===----------------------------------------------------------------------===//29// Interfaces for AllocaOp30//===----------------------------------------------------------------------===//31 32llvm::SmallVector<MemorySlot> LLVM::AllocaOp::getPromotableSlots() {33 if (!getOperation()->getBlock()->isEntryBlock())34 return {};35 36 return {MemorySlot{getResult(), getElemType()}};37}38 39Value LLVM::AllocaOp::getDefaultValue(const MemorySlot &slot,40 OpBuilder &builder) {41 return LLVM::UndefOp::create(builder, getLoc(), slot.elemType);42}43 44void LLVM::AllocaOp::handleBlockArgument(const MemorySlot &slot,45 BlockArgument argument,46 OpBuilder &builder) {47 for (Operation *user : getOperation()->getUsers())48 if (auto declareOp = llvm::dyn_cast<LLVM::DbgDeclareOp>(user))49 LLVM::DbgValueOp::create(builder, declareOp.getLoc(), argument,50 declareOp.getVarInfo(),51 declareOp.getLocationExpr());52}53 54std::optional<PromotableAllocationOpInterface>55LLVM::AllocaOp::handlePromotionComplete(const MemorySlot &slot,56 Value defaultValue,57 OpBuilder &builder) {58 if (defaultValue && defaultValue.use_empty())59 defaultValue.getDefiningOp()->erase();60 this->erase();61 return std::nullopt;62}63 64SmallVector<DestructurableMemorySlot> LLVM::AllocaOp::getDestructurableSlots() {65 if (!mlir::matchPattern(getArraySize(), m_One()))66 return {};67 68 auto destructurable = dyn_cast<DestructurableTypeInterface>(getElemType());69 if (!destructurable)70 return {};71 72 std::optional<DenseMap<Attribute, Type>> destructuredType =73 destructurable.getSubelementIndexMap();74 if (!destructuredType)75 return {};76 77 return {DestructurableMemorySlot{{getResult(), getElemType()},78 *destructuredType}};79}80 81DenseMap<Attribute, MemorySlot> LLVM::AllocaOp::destructure(82 const DestructurableMemorySlot &slot,83 const SmallPtrSetImpl<Attribute> &usedIndices, OpBuilder &builder,84 SmallVectorImpl<DestructurableAllocationOpInterface> &newAllocators) {85 assert(slot.ptr == getResult());86 builder.setInsertionPointAfter(*this);87 88 auto destructurableType = cast<DestructurableTypeInterface>(getElemType());89 DenseMap<Attribute, MemorySlot> slotMap;90 for (Attribute index : usedIndices) {91 Type elemType = destructurableType.getTypeAtIndex(index);92 assert(elemType && "used index must exist");93 auto subAlloca = LLVM::AllocaOp::create(94 builder, getLoc(), LLVM::LLVMPointerType::get(getContext()), elemType,95 getArraySize());96 newAllocators.push_back(subAlloca);97 slotMap.try_emplace<MemorySlot>(index, {subAlloca.getResult(), elemType});98 }99 100 return slotMap;101}102 103std::optional<DestructurableAllocationOpInterface>104LLVM::AllocaOp::handleDestructuringComplete(105 const DestructurableMemorySlot &slot, OpBuilder &builder) {106 assert(slot.ptr == getResult());107 this->erase();108 return std::nullopt;109}110 111//===----------------------------------------------------------------------===//112// Interfaces for LoadOp/StoreOp113//===----------------------------------------------------------------------===//114 115bool LLVM::LoadOp::loadsFrom(const MemorySlot &slot) {116 return getAddr() == slot.ptr;117}118 119bool LLVM::LoadOp::storesTo(const MemorySlot &slot) { return false; }120 121Value LLVM::LoadOp::getStored(const MemorySlot &slot, OpBuilder &builder,122 Value reachingDef, const DataLayout &dataLayout) {123 llvm_unreachable("getStored should not be called on LoadOp");124}125 126bool LLVM::StoreOp::loadsFrom(const MemorySlot &slot) { return false; }127 128bool LLVM::StoreOp::storesTo(const MemorySlot &slot) {129 return getAddr() == slot.ptr;130}131 132/// Checks if `type` can be used in any kind of conversion sequences.133static bool isSupportedTypeForConversion(Type type) {134 // Aggregate types are not bitcastable.135 if (isa<LLVM::LLVMStructType, LLVM::LLVMArrayType>(type))136 return false;137 138 if (auto vectorType = dyn_cast<VectorType>(type)) {139 // Vectors of pointers cannot be casted.140 if (isa<LLVM::LLVMPointerType>(vectorType.getElementType()))141 return false;142 // Scalable types are not supported.143 return !vectorType.isScalable();144 }145 return true;146}147 148/// Checks that `rhs` can be converted to `lhs` by a sequence of casts and149/// truncations. Checks for narrowing or widening conversion compatibility150/// depending on `narrowingConversion`.151static bool areConversionCompatible(const DataLayout &layout, Type targetType,152 Type srcType, bool narrowingConversion) {153 if (targetType == srcType)154 return true;155 156 if (!isSupportedTypeForConversion(targetType) ||157 !isSupportedTypeForConversion(srcType))158 return false;159 160 uint64_t targetSize = layout.getTypeSize(targetType);161 uint64_t srcSize = layout.getTypeSize(srcType);162 163 // Pointer casts will only be sane when the bitsize of both pointer types is164 // the same.165 if (isa<LLVM::LLVMPointerType>(targetType) &&166 isa<LLVM::LLVMPointerType>(srcType))167 return targetSize == srcSize;168 169 if (narrowingConversion)170 return targetSize <= srcSize;171 return targetSize >= srcSize;172}173 174/// Checks if `dataLayout` describes a little endian layout.175static bool isBigEndian(const DataLayout &dataLayout) {176 auto endiannessStr = dyn_cast_or_null<StringAttr>(dataLayout.getEndianness());177 return endiannessStr && endiannessStr == "big";178}179 180/// Converts a value to an integer type of the same size.181/// Assumes that the type can be converted.182static Value castToSameSizedInt(OpBuilder &builder, Location loc, Value val,183 const DataLayout &dataLayout) {184 Type type = val.getType();185 assert(isSupportedTypeForConversion(type) &&186 "expected value to have a convertible type");187 188 if (isa<IntegerType>(type))189 return val;190 191 uint64_t typeBitSize = dataLayout.getTypeSizeInBits(type);192 IntegerType valueSizeInteger = builder.getIntegerType(typeBitSize);193 194 if (isa<LLVM::LLVMPointerType>(type))195 return builder.createOrFold<LLVM::PtrToIntOp>(loc, valueSizeInteger, val);196 return builder.createOrFold<LLVM::BitcastOp>(loc, valueSizeInteger, val);197}198 199/// Converts a value with an integer type to `targetType`.200static Value castIntValueToSameSizedType(OpBuilder &builder, Location loc,201 Value val, Type targetType) {202 assert(isa<IntegerType>(val.getType()) &&203 "expected value to have an integer type");204 assert(isSupportedTypeForConversion(targetType) &&205 "expected the target type to be supported for conversions");206 if (val.getType() == targetType)207 return val;208 if (isa<LLVM::LLVMPointerType>(targetType))209 return builder.createOrFold<LLVM::IntToPtrOp>(loc, targetType, val);210 return builder.createOrFold<LLVM::BitcastOp>(loc, targetType, val);211}212 213/// Constructs operations that convert `srcValue` into a new value of type214/// `targetType`. Assumes the types have the same bitsize.215static Value castSameSizedTypes(OpBuilder &builder, Location loc,216 Value srcValue, Type targetType,217 const DataLayout &dataLayout) {218 Type srcType = srcValue.getType();219 assert(areConversionCompatible(dataLayout, targetType, srcType,220 /*narrowingConversion=*/true) &&221 "expected that the compatibility was checked before");222 223 // Nothing has to be done if the types are already the same.224 if (srcType == targetType)225 return srcValue;226 227 // In the special case of casting one pointer to another, we want to generate228 // an address space cast. Bitcasts of pointers are not allowed and using229 // pointer to integer conversions are not equivalent due to the loss of230 // provenance.231 if (isa<LLVM::LLVMPointerType>(targetType) &&232 isa<LLVM::LLVMPointerType>(srcType))233 return builder.createOrFold<LLVM::AddrSpaceCastOp>(loc, targetType,234 srcValue);235 236 // For all other castable types, casting through integers is necessary.237 Value replacement = castToSameSizedInt(builder, loc, srcValue, dataLayout);238 return castIntValueToSameSizedType(builder, loc, replacement, targetType);239}240 241/// Constructs operations that convert `srcValue` into a new value of type242/// `targetType`. Performs bit-level extraction if the source type is larger243/// than the target type. Assumes that this conversion is possible.244static Value createExtractAndCast(OpBuilder &builder, Location loc,245 Value srcValue, Type targetType,246 const DataLayout &dataLayout) {247 // Get the types of the source and target values.248 Type srcType = srcValue.getType();249 assert(areConversionCompatible(dataLayout, targetType, srcType,250 /*narrowingConversion=*/true) &&251 "expected that the compatibility was checked before");252 253 uint64_t srcTypeSize = dataLayout.getTypeSizeInBits(srcType);254 uint64_t targetTypeSize = dataLayout.getTypeSizeInBits(targetType);255 if (srcTypeSize == targetTypeSize)256 return castSameSizedTypes(builder, loc, srcValue, targetType, dataLayout);257 258 // First, cast the value to a same-sized integer type.259 Value replacement = castToSameSizedInt(builder, loc, srcValue, dataLayout);260 261 // Truncate the integer if the size of the target is less than the value.262 if (isBigEndian(dataLayout)) {263 uint64_t shiftAmount = srcTypeSize - targetTypeSize;264 auto shiftConstant = LLVM::ConstantOp::create(265 builder, loc, builder.getIntegerAttr(srcType, shiftAmount));266 replacement =267 builder.createOrFold<LLVM::LShrOp>(loc, srcValue, shiftConstant);268 }269 270 replacement = LLVM::TruncOp::create(271 builder, loc, builder.getIntegerType(targetTypeSize), replacement);272 273 // Now cast the integer to the actual target type if required.274 return castIntValueToSameSizedType(builder, loc, replacement, targetType);275}276 277/// Constructs operations that insert the bits of `srcValue` into the278/// "beginning" of `reachingDef` (beginning is endianness dependent).279/// Assumes that this conversion is possible.280static Value createInsertAndCast(OpBuilder &builder, Location loc,281 Value srcValue, Value reachingDef,282 const DataLayout &dataLayout) {283 284 assert(areConversionCompatible(dataLayout, reachingDef.getType(),285 srcValue.getType(),286 /*narrowingConversion=*/false) &&287 "expected that the compatibility was checked before");288 uint64_t valueTypeSize = dataLayout.getTypeSizeInBits(srcValue.getType());289 uint64_t slotTypeSize = dataLayout.getTypeSizeInBits(reachingDef.getType());290 if (slotTypeSize == valueTypeSize)291 return castSameSizedTypes(builder, loc, srcValue, reachingDef.getType(),292 dataLayout);293 294 // In the case where the store only overwrites parts of the memory,295 // bit fiddling is required to construct the new value.296 297 // First convert both values to integers of the same size.298 Value defAsInt = castToSameSizedInt(builder, loc, reachingDef, dataLayout);299 Value valueAsInt = castToSameSizedInt(builder, loc, srcValue, dataLayout);300 // Extend the value to the size of the reaching definition.301 valueAsInt =302 builder.createOrFold<LLVM::ZExtOp>(loc, defAsInt.getType(), valueAsInt);303 uint64_t sizeDifference = slotTypeSize - valueTypeSize;304 if (isBigEndian(dataLayout)) {305 // On big endian systems, a store to the base pointer overwrites the most306 // significant bits. To accomodate for this, the stored value needs to be307 // shifted into the according position.308 Value bigEndianShift = LLVM::ConstantOp::create(309 builder, loc,310 builder.getIntegerAttr(defAsInt.getType(), sizeDifference));311 valueAsInt =312 builder.createOrFold<LLVM::ShlOp>(loc, valueAsInt, bigEndianShift);313 }314 315 // Construct the mask that is used to erase the bits that are overwritten by316 // the store.317 APInt maskValue;318 if (isBigEndian(dataLayout)) {319 // Build a mask that has the most significant bits set to zero.320 // Note: This is the same as 2^sizeDifference - 1321 maskValue = APInt::getAllOnes(sizeDifference).zext(slotTypeSize);322 } else {323 // Build a mask that has the least significant bits set to zero.324 // Note: This is the same as -(2^valueTypeSize)325 maskValue = APInt::getAllOnes(valueTypeSize).zext(slotTypeSize);326 maskValue.flipAllBits();327 }328 329 // Mask out the affected bits ...330 Value mask = LLVM::ConstantOp::create(331 builder, loc, builder.getIntegerAttr(defAsInt.getType(), maskValue));332 Value masked = builder.createOrFold<LLVM::AndOp>(loc, defAsInt, mask);333 334 // ... and combine the result with the new value.335 Value combined = builder.createOrFold<LLVM::OrOp>(loc, masked, valueAsInt);336 337 return castIntValueToSameSizedType(builder, loc, combined,338 reachingDef.getType());339}340 341Value LLVM::StoreOp::getStored(const MemorySlot &slot, OpBuilder &builder,342 Value reachingDef,343 const DataLayout &dataLayout) {344 assert(reachingDef && reachingDef.getType() == slot.elemType &&345 "expected the reaching definition's type to match the slot's type");346 return createInsertAndCast(builder, getLoc(), getValue(), reachingDef,347 dataLayout);348}349 350bool LLVM::LoadOp::canUsesBeRemoved(351 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,352 SmallVectorImpl<OpOperand *> &newBlockingUses,353 const DataLayout &dataLayout) {354 if (blockingUses.size() != 1)355 return false;356 Value blockingUse = (*blockingUses.begin())->get();357 // If the blocking use is the slot ptr itself, there will be enough358 // context to reconstruct the result of the load at removal time, so it can359 // be removed (provided it is not volatile).360 return blockingUse == slot.ptr && getAddr() == slot.ptr &&361 areConversionCompatible(dataLayout, getResult().getType(),362 slot.elemType, /*narrowingConversion=*/true) &&363 !getVolatile_();364}365 366DeletionKind LLVM::LoadOp::removeBlockingUses(367 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,368 OpBuilder &builder, Value reachingDefinition,369 const DataLayout &dataLayout) {370 // `canUsesBeRemoved` checked this blocking use must be the loaded slot371 // pointer.372 Value newResult = createExtractAndCast(builder, getLoc(), reachingDefinition,373 getResult().getType(), dataLayout);374 getResult().replaceAllUsesWith(newResult);375 return DeletionKind::Delete;376}377 378bool LLVM::StoreOp::canUsesBeRemoved(379 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,380 SmallVectorImpl<OpOperand *> &newBlockingUses,381 const DataLayout &dataLayout) {382 if (blockingUses.size() != 1)383 return false;384 Value blockingUse = (*blockingUses.begin())->get();385 // If the blocking use is the slot ptr itself, dropping the store is386 // fine, provided we are currently promoting its target value. Don't allow a387 // store OF the slot pointer, only INTO the slot pointer.388 return blockingUse == slot.ptr && getAddr() == slot.ptr &&389 getValue() != slot.ptr &&390 areConversionCompatible(dataLayout, slot.elemType,391 getValue().getType(),392 /*narrowingConversion=*/false) &&393 !getVolatile_();394}395 396DeletionKind LLVM::StoreOp::removeBlockingUses(397 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,398 OpBuilder &builder, Value reachingDefinition,399 const DataLayout &dataLayout) {400 return DeletionKind::Delete;401}402 403/// Checks if `slot` can be accessed through the provided access type.404static bool isValidAccessType(const MemorySlot &slot, Type accessType,405 const DataLayout &dataLayout) {406 return dataLayout.getTypeSize(accessType) <=407 dataLayout.getTypeSize(slot.elemType);408}409 410LogicalResult LLVM::LoadOp::ensureOnlySafeAccesses(411 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,412 const DataLayout &dataLayout) {413 return success(getAddr() != slot.ptr ||414 isValidAccessType(slot, getType(), dataLayout));415}416 417LogicalResult LLVM::StoreOp::ensureOnlySafeAccesses(418 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,419 const DataLayout &dataLayout) {420 return success(getAddr() != slot.ptr ||421 isValidAccessType(slot, getValue().getType(), dataLayout));422}423 424/// Returns the subslot's type at the requested index.425static Type getTypeAtIndex(const DestructurableMemorySlot &slot,426 Attribute index) {427 auto subelementIndexMap =428 cast<DestructurableTypeInterface>(slot.elemType).getSubelementIndexMap();429 if (!subelementIndexMap)430 return {};431 assert(!subelementIndexMap->empty());432 433 // Note: Returns a null-type when no entry was found.434 return subelementIndexMap->lookup(index);435}436 437bool LLVM::LoadOp::canRewire(const DestructurableMemorySlot &slot,438 SmallPtrSetImpl<Attribute> &usedIndices,439 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,440 const DataLayout &dataLayout) {441 if (getVolatile_())442 return false;443 444 // A load always accesses the first element of the destructured slot.445 auto index = IntegerAttr::get(IntegerType::get(getContext(), 32), 0);446 Type subslotType = getTypeAtIndex(slot, index);447 if (!subslotType)448 return false;449 450 // The access can only be replaced when the subslot is read within its bounds.451 if (dataLayout.getTypeSize(getType()) > dataLayout.getTypeSize(subslotType))452 return false;453 454 usedIndices.insert(index);455 return true;456}457 458DeletionKind LLVM::LoadOp::rewire(const DestructurableMemorySlot &slot,459 DenseMap<Attribute, MemorySlot> &subslots,460 OpBuilder &builder,461 const DataLayout &dataLayout) {462 auto index = IntegerAttr::get(IntegerType::get(getContext(), 32), 0);463 auto it = subslots.find(index);464 assert(it != subslots.end());465 466 getAddrMutable().set(it->getSecond().ptr);467 return DeletionKind::Keep;468}469 470bool LLVM::StoreOp::canRewire(const DestructurableMemorySlot &slot,471 SmallPtrSetImpl<Attribute> &usedIndices,472 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,473 const DataLayout &dataLayout) {474 if (getVolatile_())475 return false;476 477 // Storing the pointer to memory cannot be dealt with.478 if (getValue() == slot.ptr)479 return false;480 481 // A store always accesses the first element of the destructured slot.482 auto index = IntegerAttr::get(IntegerType::get(getContext(), 32), 0);483 Type subslotType = getTypeAtIndex(slot, index);484 if (!subslotType)485 return false;486 487 // The access can only be replaced when the subslot is read within its bounds.488 if (dataLayout.getTypeSize(getValue().getType()) >489 dataLayout.getTypeSize(subslotType))490 return false;491 492 usedIndices.insert(index);493 return true;494}495 496DeletionKind LLVM::StoreOp::rewire(const DestructurableMemorySlot &slot,497 DenseMap<Attribute, MemorySlot> &subslots,498 OpBuilder &builder,499 const DataLayout &dataLayout) {500 auto index = IntegerAttr::get(IntegerType::get(getContext(), 32), 0);501 auto it = subslots.find(index);502 assert(it != subslots.end());503 504 getAddrMutable().set(it->getSecond().ptr);505 return DeletionKind::Keep;506}507 508//===----------------------------------------------------------------------===//509// Interfaces for discardable OPs510//===----------------------------------------------------------------------===//511 512/// Conditions the deletion of the operation to the removal of all its uses.513static bool forwardToUsers(Operation *op,514 SmallVectorImpl<OpOperand *> &newBlockingUses) {515 for (Value result : op->getResults())516 for (OpOperand &use : result.getUses())517 newBlockingUses.push_back(&use);518 return true;519}520 521bool LLVM::BitcastOp::canUsesBeRemoved(522 const SmallPtrSetImpl<OpOperand *> &blockingUses,523 SmallVectorImpl<OpOperand *> &newBlockingUses,524 const DataLayout &dataLayout) {525 return forwardToUsers(*this, newBlockingUses);526}527 528DeletionKind LLVM::BitcastOp::removeBlockingUses(529 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {530 return DeletionKind::Delete;531}532 533bool LLVM::AddrSpaceCastOp::canUsesBeRemoved(534 const SmallPtrSetImpl<OpOperand *> &blockingUses,535 SmallVectorImpl<OpOperand *> &newBlockingUses,536 const DataLayout &dataLayout) {537 return forwardToUsers(*this, newBlockingUses);538}539 540DeletionKind LLVM::AddrSpaceCastOp::removeBlockingUses(541 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {542 return DeletionKind::Delete;543}544 545bool LLVM::LifetimeStartOp::canUsesBeRemoved(546 const SmallPtrSetImpl<OpOperand *> &blockingUses,547 SmallVectorImpl<OpOperand *> &newBlockingUses,548 const DataLayout &dataLayout) {549 return true;550}551 552DeletionKind LLVM::LifetimeStartOp::removeBlockingUses(553 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {554 return DeletionKind::Delete;555}556 557bool LLVM::LifetimeEndOp::canUsesBeRemoved(558 const SmallPtrSetImpl<OpOperand *> &blockingUses,559 SmallVectorImpl<OpOperand *> &newBlockingUses,560 const DataLayout &dataLayout) {561 return true;562}563 564DeletionKind LLVM::LifetimeEndOp::removeBlockingUses(565 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {566 return DeletionKind::Delete;567}568 569bool LLVM::InvariantStartOp::canUsesBeRemoved(570 const SmallPtrSetImpl<OpOperand *> &blockingUses,571 SmallVectorImpl<OpOperand *> &newBlockingUses,572 const DataLayout &dataLayout) {573 return true;574}575 576DeletionKind LLVM::InvariantStartOp::removeBlockingUses(577 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {578 return DeletionKind::Delete;579}580 581bool LLVM::InvariantEndOp::canUsesBeRemoved(582 const SmallPtrSetImpl<OpOperand *> &blockingUses,583 SmallVectorImpl<OpOperand *> &newBlockingUses,584 const DataLayout &dataLayout) {585 return true;586}587 588DeletionKind LLVM::InvariantEndOp::removeBlockingUses(589 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {590 return DeletionKind::Delete;591}592 593bool LLVM::LaunderInvariantGroupOp::canUsesBeRemoved(594 const SmallPtrSetImpl<OpOperand *> &blockingUses,595 SmallVectorImpl<OpOperand *> &newBlockingUses,596 const DataLayout &dataLayout) {597 return forwardToUsers(*this, newBlockingUses);598}599 600DeletionKind LLVM::LaunderInvariantGroupOp::removeBlockingUses(601 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {602 return DeletionKind::Delete;603}604 605bool LLVM::StripInvariantGroupOp::canUsesBeRemoved(606 const SmallPtrSetImpl<OpOperand *> &blockingUses,607 SmallVectorImpl<OpOperand *> &newBlockingUses,608 const DataLayout &dataLayout) {609 return forwardToUsers(*this, newBlockingUses);610}611 612DeletionKind LLVM::StripInvariantGroupOp::removeBlockingUses(613 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {614 return DeletionKind::Delete;615}616 617bool LLVM::DbgDeclareOp::canUsesBeRemoved(618 const SmallPtrSetImpl<OpOperand *> &blockingUses,619 SmallVectorImpl<OpOperand *> &newBlockingUses,620 const DataLayout &dataLayout) {621 return true;622}623 624DeletionKind LLVM::DbgDeclareOp::removeBlockingUses(625 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {626 return DeletionKind::Delete;627}628 629bool LLVM::DbgValueOp::canUsesBeRemoved(630 const SmallPtrSetImpl<OpOperand *> &blockingUses,631 SmallVectorImpl<OpOperand *> &newBlockingUses,632 const DataLayout &dataLayout) {633 // There is only one operand that we can remove the use of.634 if (blockingUses.size() != 1)635 return false;636 637 return (*blockingUses.begin())->get() == getValue();638}639 640DeletionKind LLVM::DbgValueOp::removeBlockingUses(641 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {642 // builder by default is after '*this', but we need it before '*this'.643 builder.setInsertionPoint(*this);644 645 // Rather than dropping the debug value, replace it with undef to preserve the646 // debug local variable info. This allows the debugger to inform the user that647 // the variable has been optimized out.648 auto undef =649 UndefOp::create(builder, getValue().getLoc(), getValue().getType());650 getValueMutable().assign(undef);651 return DeletionKind::Keep;652}653 654bool LLVM::DbgDeclareOp::requiresReplacedValues() { return true; }655 656void LLVM::DbgDeclareOp::visitReplacedValues(657 ArrayRef<std::pair<Operation *, Value>> definitions, OpBuilder &builder) {658 for (auto [op, value] : definitions) {659 builder.setInsertionPointAfter(op);660 LLVM::DbgValueOp::create(builder, getLoc(), value, getVarInfo(),661 getLocationExpr());662 }663}664 665//===----------------------------------------------------------------------===//666// Interfaces for GEPOp667//===----------------------------------------------------------------------===//668 669static bool hasAllZeroIndices(LLVM::GEPOp gepOp) {670 return llvm::all_of(gepOp.getIndices(), [](auto index) {671 auto indexAttr = llvm::dyn_cast_if_present<IntegerAttr>(index);672 return indexAttr && indexAttr.getValue() == 0;673 });674}675 676bool LLVM::GEPOp::canUsesBeRemoved(677 const SmallPtrSetImpl<OpOperand *> &blockingUses,678 SmallVectorImpl<OpOperand *> &newBlockingUses,679 const DataLayout &dataLayout) {680 // GEP can be removed as long as it is a no-op and its users can be removed.681 if (!hasAllZeroIndices(*this))682 return false;683 return forwardToUsers(*this, newBlockingUses);684}685 686DeletionKind LLVM::GEPOp::removeBlockingUses(687 const SmallPtrSetImpl<OpOperand *> &blockingUses, OpBuilder &builder) {688 return DeletionKind::Delete;689}690 691/// Returns the amount of bytes the provided GEP elements will offset the692/// pointer by. Returns nullopt if no constant offset could be computed.693static std::optional<uint64_t> gepToByteOffset(const DataLayout &dataLayout,694 LLVM::GEPOp gep) {695 // Collects all indices.696 SmallVector<uint64_t> indices;697 for (auto index : gep.getIndices()) {698 auto constIndex = dyn_cast<IntegerAttr>(index);699 if (!constIndex)700 return {};701 int64_t gepIndex = constIndex.getInt();702 // Negative indices are not supported.703 if (gepIndex < 0)704 return {};705 indices.push_back(gepIndex);706 }707 708 Type currentType = gep.getElemType();709 uint64_t offset = indices[0] * dataLayout.getTypeSize(currentType);710 711 for (uint64_t index : llvm::drop_begin(indices)) {712 bool shouldCancel =713 TypeSwitch<Type, bool>(currentType)714 .Case([&](LLVM::LLVMArrayType arrayType) {715 offset +=716 index * dataLayout.getTypeSize(arrayType.getElementType());717 currentType = arrayType.getElementType();718 return false;719 })720 .Case([&](LLVM::LLVMStructType structType) {721 ArrayRef<Type> body = structType.getBody();722 assert(index < body.size() && "expected valid struct indexing");723 for (uint32_t i : llvm::seq(index)) {724 if (!structType.isPacked())725 offset = llvm::alignTo(726 offset, dataLayout.getTypeABIAlignment(body[i]));727 offset += dataLayout.getTypeSize(body[i]);728 }729 730 // Align for the current type as well.731 if (!structType.isPacked())732 offset = llvm::alignTo(733 offset, dataLayout.getTypeABIAlignment(body[index]));734 currentType = body[index];735 return false;736 })737 .Default([&](Type type) {738 LDBG() << "[sroa] Unsupported type for offset computations"739 << type;740 return true;741 });742 743 if (shouldCancel)744 return std::nullopt;745 }746 747 return offset;748}749 750namespace {751/// A struct that stores both the index into the aggregate type of the slot as752/// well as the corresponding byte offset in memory.753struct SubslotAccessInfo {754 /// The parent slot's index that the access falls into.755 uint32_t index;756 /// The offset into the subslot of the access.757 uint64_t subslotOffset;758};759} // namespace760 761/// Computes subslot access information for an access into `slot` with the given762/// offset.763/// Returns nullopt when the offset is out-of-bounds or when the access is into764/// the padding of `slot`.765static std::optional<SubslotAccessInfo>766getSubslotAccessInfo(const DestructurableMemorySlot &slot,767 const DataLayout &dataLayout, LLVM::GEPOp gep) {768 std::optional<uint64_t> offset = gepToByteOffset(dataLayout, gep);769 if (!offset)770 return {};771 772 // Helper to check that a constant index is in the bounds of the GEP index773 // representation. LLVM dialects's GEP arguments have a limited bitwidth, thus774 // this additional check is necessary.775 auto isOutOfBoundsGEPIndex = [](uint64_t index) {776 return index >= (1 << LLVM::kGEPConstantBitWidth);777 };778 779 Type type = slot.elemType;780 if (*offset >= dataLayout.getTypeSize(type))781 return {};782 return TypeSwitch<Type, std::optional<SubslotAccessInfo>>(type)783 .Case([&](LLVM::LLVMArrayType arrayType)784 -> std::optional<SubslotAccessInfo> {785 // Find which element of the array contains the offset.786 uint64_t elemSize = dataLayout.getTypeSize(arrayType.getElementType());787 uint64_t index = *offset / elemSize;788 if (isOutOfBoundsGEPIndex(index))789 return {};790 return SubslotAccessInfo{static_cast<uint32_t>(index),791 *offset - (index * elemSize)};792 })793 .Case([&](LLVM::LLVMStructType structType)794 -> std::optional<SubslotAccessInfo> {795 uint64_t distanceToStart = 0;796 // Walk over the elements of the struct to find in which of797 // them the offset is.798 for (auto [index, elem] : llvm::enumerate(structType.getBody())) {799 uint64_t elemSize = dataLayout.getTypeSize(elem);800 if (!structType.isPacked()) {801 distanceToStart = llvm::alignTo(802 distanceToStart, dataLayout.getTypeABIAlignment(elem));803 // If the offset is in padding, cancel the rewrite.804 if (offset < distanceToStart)805 return {};806 }807 808 if (offset < distanceToStart + elemSize) {809 if (isOutOfBoundsGEPIndex(index))810 return {};811 // The offset is within this element, stop iterating the812 // struct and return the index.813 return SubslotAccessInfo{static_cast<uint32_t>(index),814 *offset - distanceToStart};815 }816 817 // The offset is not within this element, continue walking818 // over the struct.819 distanceToStart += elemSize;820 }821 822 return {};823 });824}825 826/// Constructs a byte array type of the given size.827static LLVM::LLVMArrayType getByteArrayType(MLIRContext *context,828 unsigned size) {829 auto byteType = IntegerType::get(context, 8);830 return LLVM::LLVMArrayType::get(context, byteType, size);831}832 833LogicalResult LLVM::GEPOp::ensureOnlySafeAccesses(834 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,835 const DataLayout &dataLayout) {836 if (getBase() != slot.ptr)837 return success();838 std::optional<uint64_t> gepOffset = gepToByteOffset(dataLayout, *this);839 if (!gepOffset)840 return failure();841 uint64_t slotSize = dataLayout.getTypeSize(slot.elemType);842 // Check that the access is strictly inside the slot.843 if (*gepOffset >= slotSize)844 return failure();845 // Every access that remains in bounds of the remaining slot is considered846 // legal.847 mustBeSafelyUsed.emplace_back<MemorySlot>(848 {getRes(), getByteArrayType(getContext(), slotSize - *gepOffset)});849 return success();850}851 852bool LLVM::GEPOp::canRewire(const DestructurableMemorySlot &slot,853 SmallPtrSetImpl<Attribute> &usedIndices,854 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,855 const DataLayout &dataLayout) {856 if (!isa<LLVM::LLVMPointerType>(getBase().getType()))857 return false;858 859 if (getBase() != slot.ptr)860 return false;861 std::optional<SubslotAccessInfo> accessInfo =862 getSubslotAccessInfo(slot, dataLayout, *this);863 if (!accessInfo)864 return false;865 auto indexAttr =866 IntegerAttr::get(IntegerType::get(getContext(), 32), accessInfo->index);867 assert(slot.subelementTypes.contains(indexAttr));868 usedIndices.insert(indexAttr);869 870 // The remainder of the subslot should be accesses in-bounds. Thus, we create871 // a dummy slot with the size of the remainder.872 Type subslotType = slot.subelementTypes.lookup(indexAttr);873 uint64_t slotSize = dataLayout.getTypeSize(subslotType);874 LLVM::LLVMArrayType remainingSlotType =875 getByteArrayType(getContext(), slotSize - accessInfo->subslotOffset);876 mustBeSafelyUsed.emplace_back<MemorySlot>({getRes(), remainingSlotType});877 878 return true;879}880 881DeletionKind LLVM::GEPOp::rewire(const DestructurableMemorySlot &slot,882 DenseMap<Attribute, MemorySlot> &subslots,883 OpBuilder &builder,884 const DataLayout &dataLayout) {885 std::optional<SubslotAccessInfo> accessInfo =886 getSubslotAccessInfo(slot, dataLayout, *this);887 assert(accessInfo && "expected access info to be checked before");888 auto indexAttr =889 IntegerAttr::get(IntegerType::get(getContext(), 32), accessInfo->index);890 const MemorySlot &newSlot = subslots.at(indexAttr);891 892 auto byteType = IntegerType::get(builder.getContext(), 8);893 auto newPtr = builder.createOrFold<LLVM::GEPOp>(894 getLoc(), getResult().getType(), byteType, newSlot.ptr,895 ArrayRef<GEPArg>(accessInfo->subslotOffset), getNoWrapFlags());896 getResult().replaceAllUsesWith(newPtr);897 return DeletionKind::Delete;898}899 900//===----------------------------------------------------------------------===//901// Utilities for memory intrinsics902//===----------------------------------------------------------------------===//903 904namespace {905 906/// Returns the length of the given memory intrinsic in bytes if it can be known907/// at compile-time on a best-effort basis, nothing otherwise.908template <class MemIntr>909std::optional<uint64_t> getStaticMemIntrLen(MemIntr op) {910 APInt memIntrLen;911 if (!matchPattern(op.getLen(), m_ConstantInt(&memIntrLen)))912 return {};913 if (memIntrLen.getBitWidth() > 64)914 return {};915 return memIntrLen.getZExtValue();916}917 918/// Returns the length of the given memory intrinsic in bytes if it can be known919/// at compile-time on a best-effort basis, nothing otherwise.920/// Because MemcpyInlineOp has its length encoded as an attribute, this requires921/// specialized handling.922template <>923std::optional<uint64_t> getStaticMemIntrLen(LLVM::MemcpyInlineOp op) {924 APInt memIntrLen = op.getLen();925 if (memIntrLen.getBitWidth() > 64)926 return {};927 return memIntrLen.getZExtValue();928}929 930/// Returns the length of the given memory intrinsic in bytes if it can be known931/// at compile-time on a best-effort basis, nothing otherwise.932/// Because MemsetInlineOp has its length encoded as an attribute, this requires933/// specialized handling.934template <>935std::optional<uint64_t> getStaticMemIntrLen(LLVM::MemsetInlineOp op) {936 APInt memIntrLen = op.getLen();937 if (memIntrLen.getBitWidth() > 64)938 return {};939 return memIntrLen.getZExtValue();940}941 942/// Returns an integer attribute representing the length of a memset intrinsic943template <class MemsetIntr>944IntegerAttr createMemsetLenAttr(MemsetIntr op) {945 IntegerAttr memsetLenAttr;946 bool successfulMatch =947 matchPattern(op.getLen(), m_Constant<IntegerAttr>(&memsetLenAttr));948 (void)successfulMatch;949 assert(successfulMatch);950 return memsetLenAttr;951}952 953/// Returns an integer attribute representing the length of a memset intrinsic954/// Because MemsetInlineOp has its length encoded as an attribute, this requires955/// specialized handling.956template <>957IntegerAttr createMemsetLenAttr(LLVM::MemsetInlineOp op) {958 return op.getLenAttr();959}960 961/// Creates a memset intrinsic of that matches the `toReplace` intrinsic962/// using the provided parameters. There are template specializations for963/// MemsetOp and MemsetInlineOp.964template <class MemsetIntr>965void createMemsetIntr(OpBuilder &builder, MemsetIntr toReplace,966 IntegerAttr memsetLenAttr, uint64_t newMemsetSize,967 DenseMap<Attribute, MemorySlot> &subslots,968 Attribute index);969 970template <>971void createMemsetIntr(OpBuilder &builder, LLVM::MemsetOp toReplace,972 IntegerAttr memsetLenAttr, uint64_t newMemsetSize,973 DenseMap<Attribute, MemorySlot> &subslots,974 Attribute index) {975 Value newMemsetSizeValue =976 LLVM::ConstantOp::create(977 builder, toReplace.getLen().getLoc(),978 IntegerAttr::get(memsetLenAttr.getType(), newMemsetSize))979 .getResult();980 981 LLVM::MemsetOp::create(builder, toReplace.getLoc(), subslots.at(index).ptr,982 toReplace.getVal(), newMemsetSizeValue,983 toReplace.getIsVolatile());984}985 986template <>987void createMemsetIntr(OpBuilder &builder, LLVM::MemsetInlineOp toReplace,988 IntegerAttr memsetLenAttr, uint64_t newMemsetSize,989 DenseMap<Attribute, MemorySlot> &subslots,990 Attribute index) {991 auto newMemsetSizeValue =992 IntegerAttr::get(memsetLenAttr.getType(), newMemsetSize);993 994 LLVM::MemsetInlineOp::create(builder, toReplace.getLoc(),995 subslots.at(index).ptr, toReplace.getVal(),996 newMemsetSizeValue, toReplace.getIsVolatile());997}998 999} // namespace1000 1001/// Returns whether one can be sure the memory intrinsic does not write outside1002/// of the bounds of the given slot, on a best-effort basis.1003template <class MemIntr>1004static bool definitelyWritesOnlyWithinSlot(MemIntr op, const MemorySlot &slot,1005 const DataLayout &dataLayout) {1006 if (!isa<LLVM::LLVMPointerType>(slot.ptr.getType()) ||1007 op.getDst() != slot.ptr)1008 return false;1009 1010 std::optional<uint64_t> memIntrLen = getStaticMemIntrLen(op);1011 return memIntrLen && *memIntrLen <= dataLayout.getTypeSize(slot.elemType);1012}1013 1014/// Checks whether all indices are i32. This is used to check GEPs can index1015/// into them.1016static bool areAllIndicesI32(const DestructurableMemorySlot &slot) {1017 Type i32 = IntegerType::get(slot.ptr.getContext(), 32);1018 return llvm::all_of(llvm::make_first_range(slot.subelementTypes),1019 [&](Attribute index) {1020 auto intIndex = dyn_cast<IntegerAttr>(index);1021 return intIndex && intIndex.getType() == i32;1022 });1023}1024 1025//===----------------------------------------------------------------------===//1026// Interfaces for memset and memset.inline1027//===----------------------------------------------------------------------===//1028 1029template <class MemsetIntr>1030static bool memsetCanRewire(MemsetIntr op, const DestructurableMemorySlot &slot,1031 SmallPtrSetImpl<Attribute> &usedIndices,1032 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1033 const DataLayout &dataLayout) {1034 if (&slot.elemType.getDialect() != op.getOperation()->getDialect())1035 return false;1036 1037 if (op.getIsVolatile())1038 return false;1039 1040 if (!cast<DestructurableTypeInterface>(slot.elemType).getSubelementIndexMap())1041 return false;1042 1043 if (!areAllIndicesI32(slot))1044 return false;1045 1046 return definitelyWritesOnlyWithinSlot(op, slot, dataLayout);1047}1048 1049template <class MemsetIntr>1050static Value memsetGetStored(MemsetIntr op, const MemorySlot &slot,1051 OpBuilder &builder) {1052 /// Returns an integer value that is `width` bits wide representing the value1053 /// assigned to the slot by memset.1054 auto buildMemsetValue = [&](unsigned width) -> Value {1055 assert(width % 8 == 0);1056 auto intType = IntegerType::get(op.getContext(), width);1057 1058 // If we know the pattern at compile time, we can compute and assign a1059 // constant directly.1060 IntegerAttr constantPattern;1061 if (matchPattern(op.getVal(), m_Constant(&constantPattern))) {1062 assert(constantPattern.getValue().getBitWidth() == 8);1063 APInt memsetVal(/*numBits=*/width, /*val=*/0);1064 for (unsigned loBit = 0; loBit < width; loBit += 8)1065 memsetVal.insertBits(constantPattern.getValue(), loBit);1066 return LLVM::ConstantOp::create(builder, op.getLoc(),1067 IntegerAttr::get(intType, memsetVal));1068 }1069 1070 // If the output is a single byte, we can return the pattern directly.1071 if (width == 8)1072 return op.getVal();1073 1074 // Otherwise build the memset integer at runtime by repeatedly shifting the1075 // value and or-ing it with the previous value.1076 uint64_t coveredBits = 8;1077 Value currentValue =1078 LLVM::ZExtOp::create(builder, op.getLoc(), intType, op.getVal());1079 while (coveredBits < width) {1080 Value shiftBy =1081 LLVM::ConstantOp::create(builder, op.getLoc(), intType, coveredBits);1082 Value shifted =1083 LLVM::ShlOp::create(builder, op.getLoc(), currentValue, shiftBy);1084 currentValue =1085 LLVM::OrOp::create(builder, op.getLoc(), currentValue, shifted);1086 coveredBits *= 2;1087 }1088 1089 return currentValue;1090 };1091 return TypeSwitch<Type, Value>(slot.elemType)1092 .Case([&](IntegerType type) -> Value {1093 return buildMemsetValue(type.getWidth());1094 })1095 .Case([&](FloatType type) -> Value {1096 Value intVal = buildMemsetValue(type.getWidth());1097 return LLVM::BitcastOp::create(builder, op.getLoc(), type, intVal);1098 })1099 .DefaultUnreachable(1100 "getStored should not be called on memset to unsupported type");1101}1102 1103template <class MemsetIntr>1104static bool1105memsetCanUsesBeRemoved(MemsetIntr op, const MemorySlot &slot,1106 const SmallPtrSetImpl<OpOperand *> &blockingUses,1107 SmallVectorImpl<OpOperand *> &newBlockingUses,1108 const DataLayout &dataLayout) {1109 bool canConvertType =1110 TypeSwitch<Type, bool>(slot.elemType)1111 .Case<IntegerType, FloatType>([](auto type) {1112 return type.getWidth() % 8 == 0 && type.getWidth() > 0;1113 })1114 .Default(false);1115 if (!canConvertType)1116 return false;1117 1118 if (op.getIsVolatile())1119 return false;1120 1121 return getStaticMemIntrLen(op) == dataLayout.getTypeSize(slot.elemType);1122}1123 1124template <class MemsetIntr>1125static DeletionKind1126memsetRewire(MemsetIntr op, const DestructurableMemorySlot &slot,1127 DenseMap<Attribute, MemorySlot> &subslots, OpBuilder &builder,1128 const DataLayout &dataLayout) {1129 1130 std::optional<DenseMap<Attribute, Type>> types =1131 cast<DestructurableTypeInterface>(slot.elemType).getSubelementIndexMap();1132 1133 IntegerAttr memsetLenAttr = createMemsetLenAttr(op);1134 1135 bool packed = false;1136 if (auto structType = dyn_cast<LLVM::LLVMStructType>(slot.elemType))1137 packed = structType.isPacked();1138 1139 Type i32 = IntegerType::get(op.getContext(), 32);1140 uint64_t memsetLen = memsetLenAttr.getValue().getZExtValue();1141 uint64_t covered = 0;1142 for (size_t i = 0; i < types->size(); i++) {1143 // Create indices on the fly to get elements in the right order.1144 Attribute index = IntegerAttr::get(i32, i);1145 Type elemType = types->at(index);1146 uint64_t typeSize = dataLayout.getTypeSize(elemType);1147 1148 if (!packed)1149 covered =1150 llvm::alignTo(covered, dataLayout.getTypeABIAlignment(elemType));1151 1152 if (covered >= memsetLen)1153 break;1154 1155 // If this subslot is used, apply a new memset to it.1156 // Otherwise, only compute its offset within the original memset.1157 if (subslots.contains(index)) {1158 uint64_t newMemsetSize = std::min(memsetLen - covered, typeSize);1159 createMemsetIntr(builder, op, memsetLenAttr, newMemsetSize, subslots,1160 index);1161 }1162 1163 covered += typeSize;1164 }1165 1166 return DeletionKind::Delete;1167}1168 1169bool LLVM::MemsetOp::loadsFrom(const MemorySlot &slot) { return false; }1170 1171bool LLVM::MemsetOp::storesTo(const MemorySlot &slot) {1172 return getDst() == slot.ptr;1173}1174 1175Value LLVM::MemsetOp::getStored(const MemorySlot &slot, OpBuilder &builder,1176 Value reachingDef,1177 const DataLayout &dataLayout) {1178 return memsetGetStored(*this, slot, builder);1179}1180 1181bool LLVM::MemsetOp::canUsesBeRemoved(1182 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1183 SmallVectorImpl<OpOperand *> &newBlockingUses,1184 const DataLayout &dataLayout) {1185 return memsetCanUsesBeRemoved(*this, slot, blockingUses, newBlockingUses,1186 dataLayout);1187}1188 1189DeletionKind LLVM::MemsetOp::removeBlockingUses(1190 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1191 OpBuilder &builder, Value reachingDefinition,1192 const DataLayout &dataLayout) {1193 return DeletionKind::Delete;1194}1195 1196LogicalResult LLVM::MemsetOp::ensureOnlySafeAccesses(1197 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1198 const DataLayout &dataLayout) {1199 return success(definitelyWritesOnlyWithinSlot(*this, slot, dataLayout));1200}1201 1202bool LLVM::MemsetOp::canRewire(const DestructurableMemorySlot &slot,1203 SmallPtrSetImpl<Attribute> &usedIndices,1204 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1205 const DataLayout &dataLayout) {1206 return memsetCanRewire(*this, slot, usedIndices, mustBeSafelyUsed,1207 dataLayout);1208}1209 1210DeletionKind LLVM::MemsetOp::rewire(const DestructurableMemorySlot &slot,1211 DenseMap<Attribute, MemorySlot> &subslots,1212 OpBuilder &builder,1213 const DataLayout &dataLayout) {1214 return memsetRewire(*this, slot, subslots, builder, dataLayout);1215}1216 1217bool LLVM::MemsetInlineOp::loadsFrom(const MemorySlot &slot) { return false; }1218 1219bool LLVM::MemsetInlineOp::storesTo(const MemorySlot &slot) {1220 return getDst() == slot.ptr;1221}1222 1223Value LLVM::MemsetInlineOp::getStored(const MemorySlot &slot,1224 OpBuilder &builder, Value reachingDef,1225 const DataLayout &dataLayout) {1226 return memsetGetStored(*this, slot, builder);1227}1228 1229bool LLVM::MemsetInlineOp::canUsesBeRemoved(1230 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1231 SmallVectorImpl<OpOperand *> &newBlockingUses,1232 const DataLayout &dataLayout) {1233 return memsetCanUsesBeRemoved(*this, slot, blockingUses, newBlockingUses,1234 dataLayout);1235}1236 1237DeletionKind LLVM::MemsetInlineOp::removeBlockingUses(1238 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1239 OpBuilder &builder, Value reachingDefinition,1240 const DataLayout &dataLayout) {1241 return DeletionKind::Delete;1242}1243 1244LogicalResult LLVM::MemsetInlineOp::ensureOnlySafeAccesses(1245 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1246 const DataLayout &dataLayout) {1247 return success(definitelyWritesOnlyWithinSlot(*this, slot, dataLayout));1248}1249 1250bool LLVM::MemsetInlineOp::canRewire(1251 const DestructurableMemorySlot &slot,1252 SmallPtrSetImpl<Attribute> &usedIndices,1253 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1254 const DataLayout &dataLayout) {1255 return memsetCanRewire(*this, slot, usedIndices, mustBeSafelyUsed,1256 dataLayout);1257}1258 1259DeletionKind1260LLVM::MemsetInlineOp::rewire(const DestructurableMemorySlot &slot,1261 DenseMap<Attribute, MemorySlot> &subslots,1262 OpBuilder &builder, const DataLayout &dataLayout) {1263 return memsetRewire(*this, slot, subslots, builder, dataLayout);1264}1265 1266//===----------------------------------------------------------------------===//1267// Interfaces for memcpy/memmove1268//===----------------------------------------------------------------------===//1269 1270template <class MemcpyLike>1271static bool memcpyLoadsFrom(MemcpyLike op, const MemorySlot &slot) {1272 return op.getSrc() == slot.ptr;1273}1274 1275template <class MemcpyLike>1276static bool memcpyStoresTo(MemcpyLike op, const MemorySlot &slot) {1277 return op.getDst() == slot.ptr;1278}1279 1280template <class MemcpyLike>1281static Value memcpyGetStored(MemcpyLike op, const MemorySlot &slot,1282 OpBuilder &builder) {1283 return LLVM::LoadOp::create(builder, op.getLoc(), slot.elemType, op.getSrc());1284}1285 1286template <class MemcpyLike>1287static bool1288memcpyCanUsesBeRemoved(MemcpyLike op, const MemorySlot &slot,1289 const SmallPtrSetImpl<OpOperand *> &blockingUses,1290 SmallVectorImpl<OpOperand *> &newBlockingUses,1291 const DataLayout &dataLayout) {1292 // If source and destination are the same, memcpy behavior is undefined and1293 // memmove is a no-op. Because there is no memory change happening here,1294 // simplifying such operations is left to canonicalization.1295 if (op.getDst() == op.getSrc())1296 return false;1297 1298 if (op.getIsVolatile())1299 return false;1300 1301 return getStaticMemIntrLen(op) == dataLayout.getTypeSize(slot.elemType);1302}1303 1304template <class MemcpyLike>1305static DeletionKind1306memcpyRemoveBlockingUses(MemcpyLike op, const MemorySlot &slot,1307 const SmallPtrSetImpl<OpOperand *> &blockingUses,1308 OpBuilder &builder, Value reachingDefinition) {1309 if (op.loadsFrom(slot))1310 LLVM::StoreOp::create(builder, op.getLoc(), reachingDefinition,1311 op.getDst());1312 return DeletionKind::Delete;1313}1314 1315template <class MemcpyLike>1316static LogicalResult1317memcpyEnsureOnlySafeAccesses(MemcpyLike op, const MemorySlot &slot,1318 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed) {1319 DataLayout dataLayout = DataLayout::closest(op);1320 // While rewiring memcpy-like intrinsics only supports full copies, partial1321 // copies are still safe accesses so it is enough to only check for writes1322 // within bounds.1323 return success(definitelyWritesOnlyWithinSlot(op, slot, dataLayout));1324}1325 1326template <class MemcpyLike>1327static bool memcpyCanRewire(MemcpyLike op, const DestructurableMemorySlot &slot,1328 SmallPtrSetImpl<Attribute> &usedIndices,1329 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1330 const DataLayout &dataLayout) {1331 if (op.getIsVolatile())1332 return false;1333 1334 if (!cast<DestructurableTypeInterface>(slot.elemType).getSubelementIndexMap())1335 return false;1336 1337 if (!areAllIndicesI32(slot))1338 return false;1339 1340 // Only full copies are supported.1341 if (getStaticMemIntrLen(op) != dataLayout.getTypeSize(slot.elemType))1342 return false;1343 1344 if (op.getSrc() == slot.ptr)1345 usedIndices.insert_range(llvm::make_first_range(slot.subelementTypes));1346 1347 return true;1348}1349 1350namespace {1351 1352template <class MemcpyLike>1353void createMemcpyLikeToReplace(OpBuilder &builder, const DataLayout &layout,1354 MemcpyLike toReplace, Value dst, Value src,1355 Type toCpy, bool isVolatile) {1356 Value memcpySize =1357 LLVM::ConstantOp::create(builder, toReplace.getLoc(),1358 IntegerAttr::get(toReplace.getLen().getType(),1359 layout.getTypeSize(toCpy)));1360 MemcpyLike::create(builder, toReplace.getLoc(), dst, src, memcpySize,1361 isVolatile);1362}1363 1364template <>1365void createMemcpyLikeToReplace(OpBuilder &builder, const DataLayout &layout,1366 LLVM::MemcpyInlineOp toReplace, Value dst,1367 Value src, Type toCpy, bool isVolatile) {1368 Type lenType = IntegerType::get(toReplace->getContext(),1369 toReplace.getLen().getBitWidth());1370 LLVM::MemcpyInlineOp::create(1371 builder, toReplace.getLoc(), dst, src,1372 IntegerAttr::get(lenType, layout.getTypeSize(toCpy)), isVolatile);1373}1374 1375} // namespace1376 1377/// Rewires a memcpy-like operation. Only copies to or from the full slot are1378/// supported.1379template <class MemcpyLike>1380static DeletionKind1381memcpyRewire(MemcpyLike op, const DestructurableMemorySlot &slot,1382 DenseMap<Attribute, MemorySlot> &subslots, OpBuilder &builder,1383 const DataLayout &dataLayout) {1384 if (subslots.empty())1385 return DeletionKind::Delete;1386 1387 assert((slot.ptr == op.getDst()) != (slot.ptr == op.getSrc()));1388 bool isDst = slot.ptr == op.getDst();1389 1390#ifndef NDEBUG1391 size_t slotsTreated = 0;1392#endif1393 1394 // It was previously checked that index types are consistent, so this type can1395 // be fetched now.1396 Type indexType = cast<IntegerAttr>(subslots.begin()->first).getType();1397 for (size_t i = 0, e = slot.subelementTypes.size(); i != e; i++) {1398 Attribute index = IntegerAttr::get(indexType, i);1399 if (!subslots.contains(index))1400 continue;1401 const MemorySlot &subslot = subslots.at(index);1402 1403#ifndef NDEBUG1404 slotsTreated++;1405#endif1406 1407 // First get a pointer to the equivalent of this subslot from the source1408 // pointer.1409 SmallVector<LLVM::GEPArg> gepIndices{1410 0, static_cast<int32_t>(1411 cast<IntegerAttr>(index).getValue().getZExtValue())};1412 Value subslotPtrInOther = LLVM::GEPOp::create(1413 builder, op.getLoc(), LLVM::LLVMPointerType::get(op.getContext()),1414 slot.elemType, isDst ? op.getSrc() : op.getDst(), gepIndices);1415 1416 // Then create a new memcpy out of this source pointer.1417 createMemcpyLikeToReplace(builder, dataLayout, op,1418 isDst ? subslot.ptr : subslotPtrInOther,1419 isDst ? subslotPtrInOther : subslot.ptr,1420 subslot.elemType, op.getIsVolatile());1421 }1422 1423 assert(subslots.size() == slotsTreated);1424 1425 return DeletionKind::Delete;1426}1427 1428bool LLVM::MemcpyOp::loadsFrom(const MemorySlot &slot) {1429 return memcpyLoadsFrom(*this, slot);1430}1431 1432bool LLVM::MemcpyOp::storesTo(const MemorySlot &slot) {1433 return memcpyStoresTo(*this, slot);1434}1435 1436Value LLVM::MemcpyOp::getStored(const MemorySlot &slot, OpBuilder &builder,1437 Value reachingDef,1438 const DataLayout &dataLayout) {1439 return memcpyGetStored(*this, slot, builder);1440}1441 1442bool LLVM::MemcpyOp::canUsesBeRemoved(1443 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1444 SmallVectorImpl<OpOperand *> &newBlockingUses,1445 const DataLayout &dataLayout) {1446 return memcpyCanUsesBeRemoved(*this, slot, blockingUses, newBlockingUses,1447 dataLayout);1448}1449 1450DeletionKind LLVM::MemcpyOp::removeBlockingUses(1451 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1452 OpBuilder &builder, Value reachingDefinition,1453 const DataLayout &dataLayout) {1454 return memcpyRemoveBlockingUses(*this, slot, blockingUses, builder,1455 reachingDefinition);1456}1457 1458LogicalResult LLVM::MemcpyOp::ensureOnlySafeAccesses(1459 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1460 const DataLayout &dataLayout) {1461 return memcpyEnsureOnlySafeAccesses(*this, slot, mustBeSafelyUsed);1462}1463 1464bool LLVM::MemcpyOp::canRewire(const DestructurableMemorySlot &slot,1465 SmallPtrSetImpl<Attribute> &usedIndices,1466 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1467 const DataLayout &dataLayout) {1468 return memcpyCanRewire(*this, slot, usedIndices, mustBeSafelyUsed,1469 dataLayout);1470}1471 1472DeletionKind LLVM::MemcpyOp::rewire(const DestructurableMemorySlot &slot,1473 DenseMap<Attribute, MemorySlot> &subslots,1474 OpBuilder &builder,1475 const DataLayout &dataLayout) {1476 return memcpyRewire(*this, slot, subslots, builder, dataLayout);1477}1478 1479bool LLVM::MemcpyInlineOp::loadsFrom(const MemorySlot &slot) {1480 return memcpyLoadsFrom(*this, slot);1481}1482 1483bool LLVM::MemcpyInlineOp::storesTo(const MemorySlot &slot) {1484 return memcpyStoresTo(*this, slot);1485}1486 1487Value LLVM::MemcpyInlineOp::getStored(const MemorySlot &slot,1488 OpBuilder &builder, Value reachingDef,1489 const DataLayout &dataLayout) {1490 return memcpyGetStored(*this, slot, builder);1491}1492 1493bool LLVM::MemcpyInlineOp::canUsesBeRemoved(1494 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1495 SmallVectorImpl<OpOperand *> &newBlockingUses,1496 const DataLayout &dataLayout) {1497 return memcpyCanUsesBeRemoved(*this, slot, blockingUses, newBlockingUses,1498 dataLayout);1499}1500 1501DeletionKind LLVM::MemcpyInlineOp::removeBlockingUses(1502 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1503 OpBuilder &builder, Value reachingDefinition,1504 const DataLayout &dataLayout) {1505 return memcpyRemoveBlockingUses(*this, slot, blockingUses, builder,1506 reachingDefinition);1507}1508 1509LogicalResult LLVM::MemcpyInlineOp::ensureOnlySafeAccesses(1510 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1511 const DataLayout &dataLayout) {1512 return memcpyEnsureOnlySafeAccesses(*this, slot, mustBeSafelyUsed);1513}1514 1515bool LLVM::MemcpyInlineOp::canRewire(1516 const DestructurableMemorySlot &slot,1517 SmallPtrSetImpl<Attribute> &usedIndices,1518 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1519 const DataLayout &dataLayout) {1520 return memcpyCanRewire(*this, slot, usedIndices, mustBeSafelyUsed,1521 dataLayout);1522}1523 1524DeletionKind1525LLVM::MemcpyInlineOp::rewire(const DestructurableMemorySlot &slot,1526 DenseMap<Attribute, MemorySlot> &subslots,1527 OpBuilder &builder, const DataLayout &dataLayout) {1528 return memcpyRewire(*this, slot, subslots, builder, dataLayout);1529}1530 1531bool LLVM::MemmoveOp::loadsFrom(const MemorySlot &slot) {1532 return memcpyLoadsFrom(*this, slot);1533}1534 1535bool LLVM::MemmoveOp::storesTo(const MemorySlot &slot) {1536 return memcpyStoresTo(*this, slot);1537}1538 1539Value LLVM::MemmoveOp::getStored(const MemorySlot &slot, OpBuilder &builder,1540 Value reachingDef,1541 const DataLayout &dataLayout) {1542 return memcpyGetStored(*this, slot, builder);1543}1544 1545bool LLVM::MemmoveOp::canUsesBeRemoved(1546 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1547 SmallVectorImpl<OpOperand *> &newBlockingUses,1548 const DataLayout &dataLayout) {1549 return memcpyCanUsesBeRemoved(*this, slot, blockingUses, newBlockingUses,1550 dataLayout);1551}1552 1553DeletionKind LLVM::MemmoveOp::removeBlockingUses(1554 const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,1555 OpBuilder &builder, Value reachingDefinition,1556 const DataLayout &dataLayout) {1557 return memcpyRemoveBlockingUses(*this, slot, blockingUses, builder,1558 reachingDefinition);1559}1560 1561LogicalResult LLVM::MemmoveOp::ensureOnlySafeAccesses(1562 const MemorySlot &slot, SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1563 const DataLayout &dataLayout) {1564 return memcpyEnsureOnlySafeAccesses(*this, slot, mustBeSafelyUsed);1565}1566 1567bool LLVM::MemmoveOp::canRewire(const DestructurableMemorySlot &slot,1568 SmallPtrSetImpl<Attribute> &usedIndices,1569 SmallVectorImpl<MemorySlot> &mustBeSafelyUsed,1570 const DataLayout &dataLayout) {1571 return memcpyCanRewire(*this, slot, usedIndices, mustBeSafelyUsed,1572 dataLayout);1573}1574 1575DeletionKind LLVM::MemmoveOp::rewire(const DestructurableMemorySlot &slot,1576 DenseMap<Attribute, MemorySlot> &subslots,1577 OpBuilder &builder,1578 const DataLayout &dataLayout) {1579 return memcpyRewire(*this, slot, subslots, builder, dataLayout);1580}1581 1582//===----------------------------------------------------------------------===//1583// Interfaces for destructurable types1584//===----------------------------------------------------------------------===//1585 1586std::optional<DenseMap<Attribute, Type>>1587LLVM::LLVMStructType::getSubelementIndexMap() const {1588 Type i32 = IntegerType::get(getContext(), 32);1589 DenseMap<Attribute, Type> destructured;1590 for (const auto &[index, elemType] : llvm::enumerate(getBody()))1591 destructured.insert({IntegerAttr::get(i32, index), elemType});1592 return destructured;1593}1594 1595Type LLVM::LLVMStructType::getTypeAtIndex(Attribute index) const {1596 auto indexAttr = llvm::dyn_cast<IntegerAttr>(index);1597 if (!indexAttr || !indexAttr.getType().isInteger(32))1598 return {};1599 int32_t indexInt = indexAttr.getInt();1600 ArrayRef<Type> body = getBody();1601 if (indexInt < 0 || body.size() <= static_cast<uint32_t>(indexInt))1602 return {};1603 return body[indexInt];1604}1605 1606std::optional<DenseMap<Attribute, Type>>1607LLVM::LLVMArrayType::getSubelementIndexMap() const {1608 constexpr size_t maxArraySizeForDestructuring = 16;1609 if (getNumElements() > maxArraySizeForDestructuring)1610 return {};1611 int32_t numElements = getNumElements();1612 1613 Type i32 = IntegerType::get(getContext(), 32);1614 DenseMap<Attribute, Type> destructured;1615 for (int32_t index = 0; index < numElements; ++index)1616 destructured.insert({IntegerAttr::get(i32, index), getElementType()});1617 return destructured;1618}1619 1620Type LLVM::LLVMArrayType::getTypeAtIndex(Attribute index) const {1621 auto indexAttr = llvm::dyn_cast<IntegerAttr>(index);1622 if (!indexAttr || !indexAttr.getType().isInteger(32))1623 return {};1624 int32_t indexInt = indexAttr.getInt();1625 if (indexInt < 0 || getNumElements() <= static_cast<uint32_t>(indexInt))1626 return {};1627 return getElementType();1628}1629