brintos

brintos / llvm-project-archived public Read only

0
0
Text · 11.8 KiB · 4dfba74 Raw
317 lines · cpp
1//===- BufferDeallocationOpInterface.cpp ----------------------------------===//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#include "mlir/Dialect/Bufferization/IR/BufferDeallocationOpInterface.h"10#include "mlir/Dialect/Bufferization/IR/Bufferization.h"11#include "mlir/Dialect/MemRef/IR/MemRef.h"12#include "mlir/IR/AsmState.h"13#include "mlir/IR/Operation.h"14#include "mlir/IR/TypeUtilities.h"15#include "mlir/IR/Value.h"16#include "llvm/ADT/SetOperations.h"17 18//===----------------------------------------------------------------------===//19// BufferDeallocationOpInterface20//===----------------------------------------------------------------------===//21 22namespace mlir {23namespace bufferization {24 25#include "mlir/Dialect/Bufferization/IR/BufferDeallocationOpInterface.cpp.inc"26 27} // namespace bufferization28} // namespace mlir29 30using namespace mlir;31using namespace bufferization;32 33//===----------------------------------------------------------------------===//34// Helpers35//===----------------------------------------------------------------------===//36 37static Value buildBoolValue(OpBuilder &builder, Location loc, bool value) {38  return arith::ConstantOp::create(builder, loc, builder.getBoolAttr(value));39}40 41static bool isMemref(Value v) { return isa<BaseMemRefType>(v.getType()); }42 43//===----------------------------------------------------------------------===//44// Ownership45//===----------------------------------------------------------------------===//46 47Ownership::Ownership(Value indicator)48    : indicator(indicator), state(State::Unique) {}49 50Ownership Ownership::getUnknown() {51  Ownership unknown;52  unknown.indicator = Value();53  unknown.state = State::Unknown;54  return unknown;55}56Ownership Ownership::getUnique(Value indicator) { return Ownership(indicator); }57Ownership Ownership::getUninitialized() { return Ownership(); }58 59bool Ownership::isUninitialized() const {60  return state == State::Uninitialized;61}62bool Ownership::isUnique() const { return state == State::Unique; }63bool Ownership::isUnknown() const { return state == State::Unknown; }64 65Value Ownership::getIndicator() const {66  assert(isUnique() && "must have unique ownership to get the indicator");67  return indicator;68}69 70Ownership Ownership::getCombined(Ownership other) const {71  if (other.isUninitialized())72    return *this;73  if (isUninitialized())74    return other;75 76  if (!isUnique() || !other.isUnique())77    return getUnknown();78 79  // Since we create a new constant i1 value for (almost) each use-site, we80  // should compare the actual value rather than just the SSA Value to avoid81  // unnecessary invalidations.82  if (isEqualConstantIntOrValue(indicator, other.indicator))83    return *this;84 85  // Return the join of the lattice if the indicator of both ownerships cannot86  // be merged.87  return getUnknown();88}89 90void Ownership::combine(Ownership other) { *this = getCombined(other); }91 92//===----------------------------------------------------------------------===//93// DeallocationState94//===----------------------------------------------------------------------===//95 96DeallocationState::DeallocationState(Operation *op,97                                     SymbolTableCollection &symbolTables)98    : symbolTable(symbolTables), liveness(op) {}99 100void DeallocationState::updateOwnership(Value memref, Ownership ownership,101                                        Block *block) {102  // In most cases we care about the block where the value is defined.103  if (block == nullptr)104    block = memref.getParentBlock();105 106  // Update ownership of current memref itself.107  ownershipMap[{memref, block}].combine(ownership);108}109 110void DeallocationState::resetOwnerships(ValueRange memrefs, Block *block) {111  for (Value val : memrefs)112    ownershipMap[{val, block}] = Ownership::getUninitialized();113}114 115Ownership DeallocationState::getOwnership(Value memref, Block *block) const {116  return ownershipMap.lookup({memref, block});117}118 119void DeallocationState::addMemrefToDeallocate(Value memref, Block *block) {120  memrefsToDeallocatePerBlock[block].push_back(memref);121}122 123void DeallocationState::dropMemrefToDeallocate(Value memref, Block *block) {124  llvm::erase(memrefsToDeallocatePerBlock[block], memref);125}126 127void DeallocationState::getLiveMemrefsIn(Block *block,128                                         SmallVectorImpl<Value> &memrefs) {129  SmallVector<Value> liveMemrefs(130      llvm::make_filter_range(liveness.getLiveIn(block), isMemref));131  llvm::sort(liveMemrefs, ValueComparator());132  memrefs.append(liveMemrefs);133}134 135std::pair<Value, Value>136DeallocationState::getMemrefWithUniqueOwnership(OpBuilder &builder,137                                                Value memref, Block *block) {138  auto iter = ownershipMap.find({memref, block});139  assert(iter != ownershipMap.end() &&140         "Value must already have been registered in the ownership map");141 142  Ownership ownership = iter->second;143  if (ownership.isUnique())144    return {memref, ownership.getIndicator()};145 146  // Instead of inserting a clone operation we could also insert a dealloc147  // operation earlier in the block and use the updated ownerships returned by148  // the op for the retained values. Alternatively, we could insert code to149  // check aliasing at runtime and use this information to combine two unique150  // ownerships more intelligently to not end up with an 'Unknown' ownership in151  // the first place.152  auto cloneOp =153      bufferization::CloneOp::create(builder, memref.getLoc(), memref);154  Value condition = buildBoolValue(builder, memref.getLoc(), true);155  Value newMemref = cloneOp.getResult();156  updateOwnership(newMemref, condition);157  memrefsToDeallocatePerBlock[newMemref.getParentBlock()].push_back(newMemref);158  return {newMemref, condition};159}160 161void DeallocationState::getMemrefsToRetain(162    Block *fromBlock, Block *toBlock, ValueRange destOperands,163    SmallVectorImpl<Value> &toRetain) const {164  for (Value operand : destOperands) {165    if (!isMemref(operand))166      continue;167    toRetain.push_back(operand);168  }169 170  SmallPtrSet<Value, 16> liveOut;171  for (auto val : liveness.getLiveOut(fromBlock))172    if (isMemref(val))173      liveOut.insert(val);174 175  if (toBlock)176    llvm::set_intersect(liveOut, liveness.getLiveIn(toBlock));177 178  // liveOut has non-deterministic order because it was constructed by iterating179  // over a hash-set.180  SmallVector<Value> retainedByLiveness(liveOut.begin(), liveOut.end());181  llvm::sort(retainedByLiveness, ValueComparator());182  toRetain.append(retainedByLiveness);183}184 185LogicalResult DeallocationState::getMemrefsAndConditionsToDeallocate(186    OpBuilder &builder, Location loc, Block *block,187    SmallVectorImpl<Value> &memrefs, SmallVectorImpl<Value> &conditions) const {188 189  for (auto [i, memref] :190       llvm::enumerate(memrefsToDeallocatePerBlock.lookup(block))) {191    Ownership ownership = ownershipMap.lookup({memref, block});192    if (!ownership.isUnique())193      return emitError(memref.getLoc(),194                       "MemRef value does not have valid ownership");195 196    // Simply cast unranked MemRefs to ranked memrefs with 0 dimensions such197    // that we can call extract_strided_metadata on it.198    if (auto unrankedMemRefTy = dyn_cast<UnrankedMemRefType>(memref.getType()))199      memref = memref::ReinterpretCastOp::create(200          builder, loc, memref,201          /*offset=*/builder.getIndexAttr(0),202          /*sizes=*/ArrayRef<OpFoldResult>{},203          /*strides=*/ArrayRef<OpFoldResult>{});204 205    // Use the `memref.extract_strided_metadata` operation to get the base206    // memref. This is needed because the same MemRef that was produced by the207    // alloc operation has to be passed to the dealloc operation. Passing208    // subviews, etc. to a dealloc operation is not allowed.209    memrefs.push_back(210        memref::ExtractStridedMetadataOp::create(builder, loc, memref)211            .getResult(0));212    conditions.push_back(ownership.getIndicator());213  }214 215  return success();216}217 218//===----------------------------------------------------------------------===//219// ValueComparator220//===----------------------------------------------------------------------===//221 222bool ValueComparator::operator()(const Value &lhs, const Value &rhs) const {223  if (lhs == rhs)224    return false;225 226  // Block arguments are less than results.227  bool lhsIsBBArg = isa<BlockArgument>(lhs);228  if (lhsIsBBArg != isa<BlockArgument>(rhs)) {229    return lhsIsBBArg;230  }231 232  Region *lhsRegion;233  Region *rhsRegion;234  if (lhsIsBBArg) {235    auto lhsBBArg = llvm::cast<BlockArgument>(lhs);236    auto rhsBBArg = llvm::cast<BlockArgument>(rhs);237    if (lhsBBArg.getArgNumber() != rhsBBArg.getArgNumber()) {238      return lhsBBArg.getArgNumber() < rhsBBArg.getArgNumber();239    }240    lhsRegion = lhsBBArg.getParentRegion();241    rhsRegion = rhsBBArg.getParentRegion();242    assert(lhsRegion != rhsRegion &&243           "lhsRegion == rhsRegion implies lhs == rhs");244  } else if (lhs.getDefiningOp() == rhs.getDefiningOp()) {245    return llvm::cast<OpResult>(lhs).getResultNumber() <246           llvm::cast<OpResult>(rhs).getResultNumber();247  } else {248    lhsRegion = lhs.getDefiningOp()->getParentRegion();249    rhsRegion = rhs.getDefiningOp()->getParentRegion();250    if (lhsRegion == rhsRegion) {251      return lhs.getDefiningOp()->isBeforeInBlock(rhs.getDefiningOp());252    }253  }254 255  // lhsRegion != rhsRegion, so if we look at their ancestor chain, they256  // - have different heights257  // - or there's a spot where their region numbers differ258  // - or their parent regions are the same and their parent ops are259  //   different.260  while (lhsRegion && rhsRegion) {261    if (lhsRegion->getRegionNumber() != rhsRegion->getRegionNumber()) {262      return lhsRegion->getRegionNumber() < rhsRegion->getRegionNumber();263    }264    if (lhsRegion->getParentRegion() == rhsRegion->getParentRegion()) {265      return lhsRegion->getParentOp()->isBeforeInBlock(266          rhsRegion->getParentOp());267    }268    lhsRegion = lhsRegion->getParentRegion();269    rhsRegion = rhsRegion->getParentRegion();270  }271  if (rhsRegion)272    return true;273  assert(lhsRegion && "this should only happen if lhs == rhs");274  return false;275}276 277//===----------------------------------------------------------------------===//278// Implementation utilities279//===----------------------------------------------------------------------===//280 281FailureOr<Operation *> deallocation_impl::insertDeallocOpForReturnLike(282    DeallocationState &state, Operation *op, ValueRange operands,283    SmallVectorImpl<Value> &updatedOperandOwnerships) {284  assert(op->hasTrait<OpTrait::IsTerminator>() && "must be a terminator");285  assert(!op->hasSuccessors() && "must not have any successors");286  // Collect the values to deallocate and retain and use them to create the287  // dealloc operation.288  OpBuilder builder(op);289  Block *block = op->getBlock();290  SmallVector<Value> memrefs, conditions, toRetain;291  if (failed(state.getMemrefsAndConditionsToDeallocate(292          builder, op->getLoc(), block, memrefs, conditions)))293    return failure();294 295  state.getMemrefsToRetain(block, /*toBlock=*/nullptr, operands, toRetain);296  if (memrefs.empty() && toRetain.empty())297    return op;298 299  auto deallocOp = bufferization::DeallocOp::create(300      builder, op->getLoc(), memrefs, conditions, toRetain);301 302  // We want to replace the current ownership of the retained values with the303  // result values of the dealloc operation as they are always unique.304  state.resetOwnerships(deallocOp.getRetained(), block);305  for (auto [retained, ownership] :306       llvm::zip(deallocOp.getRetained(), deallocOp.getUpdatedConditions()))307    state.updateOwnership(retained, ownership, block);308 309  unsigned numMemrefOperands = llvm::count_if(operands, isMemref);310  auto newOperandOwnerships =311      deallocOp.getUpdatedConditions().take_front(numMemrefOperands);312  updatedOperandOwnerships.append(newOperandOwnerships.begin(),313                                  newOperandOwnerships.end());314 315  return op;316}317