brintos

brintos / llvm-project-archived public Read only

0
0
Text · 22.5 KiB · 522e914 Raw
536 lines · cpp
1//===- MemRefBuilder.cpp - Helper for LLVM MemRef equivalents -------------===//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/Conversion/LLVMCommon/MemRefBuilder.h"10#include "MemRefDescriptor.h"11#include "mlir/Conversion/LLVMCommon/TypeConverter.h"12#include "mlir/Dialect/LLVMIR/LLVMDialect.h"13#include "mlir/Dialect/LLVMIR/LLVMTypes.h"14#include "mlir/IR/Builders.h"15#include "llvm/Support/MathExtras.h"16 17using namespace mlir;18 19//===----------------------------------------------------------------------===//20// MemRefDescriptor implementation21//===----------------------------------------------------------------------===//22 23/// Construct a helper for the given descriptor value.24MemRefDescriptor::MemRefDescriptor(Value descriptor)25    : StructBuilder(descriptor) {26  assert(value != nullptr && "value cannot be null");27  indexType = cast<LLVM::LLVMStructType>(value.getType())28                  .getBody()[kOffsetPosInMemRefDescriptor];29}30 31/// Builds IR creating an `undef` value of the descriptor type.32MemRefDescriptor MemRefDescriptor::poison(OpBuilder &builder, Location loc,33                                          Type descriptorType) {34 35  Value descriptor = LLVM::PoisonOp::create(builder, loc, descriptorType);36  return MemRefDescriptor(descriptor);37}38 39/// Builds IR creating a MemRef descriptor that represents `type` and40/// populates it with static shape and stride information extracted from the41/// type.42MemRefDescriptor43MemRefDescriptor::fromStaticShape(OpBuilder &builder, Location loc,44                                  const LLVMTypeConverter &typeConverter,45                                  MemRefType type, Value memory) {46  return fromStaticShape(builder, loc, typeConverter, type, memory, memory);47}48 49MemRefDescriptor MemRefDescriptor::fromStaticShape(50    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,51    MemRefType type, Value memory, Value alignedMemory) {52  assert(type.hasStaticShape() && "unexpected dynamic shape");53 54  // Extract all strides and offsets and verify they are static.55  auto [strides, offset] = type.getStridesAndOffset();56  assert(ShapedType::isStatic(offset) && "expected static offset");57  assert(!llvm::any_of(strides, ShapedType::isDynamic) &&58         "expected static strides");59 60  auto convertedType = typeConverter.convertType(type);61  assert(convertedType && "unexpected failure in memref type conversion");62 63  auto descr = MemRefDescriptor::poison(builder, loc, convertedType);64  descr.setAllocatedPtr(builder, loc, memory);65  descr.setAlignedPtr(builder, loc, alignedMemory);66  descr.setConstantOffset(builder, loc, offset);67 68  // Fill in sizes and strides69  for (unsigned i = 0, e = type.getRank(); i != e; ++i) {70    descr.setConstantSize(builder, loc, i, type.getDimSize(i));71    descr.setConstantStride(builder, loc, i, strides[i]);72  }73  return descr;74}75 76/// Builds IR extracting the allocated pointer from the descriptor.77Value MemRefDescriptor::allocatedPtr(OpBuilder &builder, Location loc) {78  return extractPtr(builder, loc, kAllocatedPtrPosInMemRefDescriptor);79}80 81/// Builds IR inserting the allocated pointer into the descriptor.82void MemRefDescriptor::setAllocatedPtr(OpBuilder &builder, Location loc,83                                       Value ptr) {84  setPtr(builder, loc, kAllocatedPtrPosInMemRefDescriptor, ptr);85}86 87/// Builds IR extracting the aligned pointer from the descriptor.88Value MemRefDescriptor::alignedPtr(OpBuilder &builder, Location loc) {89  return extractPtr(builder, loc, kAlignedPtrPosInMemRefDescriptor);90}91 92/// Builds IR inserting the aligned pointer into the descriptor.93void MemRefDescriptor::setAlignedPtr(OpBuilder &builder, Location loc,94                                     Value ptr) {95  setPtr(builder, loc, kAlignedPtrPosInMemRefDescriptor, ptr);96}97 98// Creates a constant Op producing a value of `resultType` from an index-typed99// integer attribute.100static Value createIndexAttrConstant(OpBuilder &builder, Location loc,101                                     Type resultType, int64_t value) {102  return LLVM::ConstantOp::create(builder, loc, resultType,103                                  builder.getIndexAttr(value));104}105 106/// Builds IR extracting the offset from the descriptor.107Value MemRefDescriptor::offset(OpBuilder &builder, Location loc) {108  return LLVM::ExtractValueOp::create(builder, loc, value,109                                      kOffsetPosInMemRefDescriptor);110}111 112/// Builds IR inserting the offset into the descriptor.113void MemRefDescriptor::setOffset(OpBuilder &builder, Location loc,114                                 Value offset) {115  value = LLVM::InsertValueOp::create(builder, loc, value, offset,116                                      kOffsetPosInMemRefDescriptor);117}118 119/// Builds IR inserting the offset into the descriptor.120void MemRefDescriptor::setConstantOffset(OpBuilder &builder, Location loc,121                                         uint64_t offset) {122  setOffset(builder, loc,123            createIndexAttrConstant(builder, loc, indexType, offset));124}125 126/// Builds IR extracting the pos-th size from the descriptor.127Value MemRefDescriptor::size(OpBuilder &builder, Location loc, unsigned pos) {128  return LLVM::ExtractValueOp::create(129      builder, loc, value,130      ArrayRef<int64_t>({kSizePosInMemRefDescriptor, pos}));131}132 133Value MemRefDescriptor::size(OpBuilder &builder, Location loc, Value pos,134                             int64_t rank) {135  auto arrayTy = LLVM::LLVMArrayType::get(indexType, rank);136 137  auto ptrTy = LLVM::LLVMPointerType::get(builder.getContext());138 139  // Copy size values to stack-allocated memory.140  auto one = createIndexAttrConstant(builder, loc, indexType, 1);141  auto sizes = LLVM::ExtractValueOp::create(142      builder, loc, value,143      llvm::ArrayRef<int64_t>({kSizePosInMemRefDescriptor}));144  auto sizesPtr = LLVM::AllocaOp::create(builder, loc, ptrTy, arrayTy, one,145                                         /*alignment=*/0);146  LLVM::StoreOp::create(builder, loc, sizes, sizesPtr);147 148  // Load an return size value of interest.149  auto resultPtr = LLVM::GEPOp::create(builder, loc, ptrTy, arrayTy, sizesPtr,150                                       ArrayRef<LLVM::GEPArg>{0, pos});151  return LLVM::LoadOp::create(builder, loc, indexType, resultPtr);152}153 154/// Builds IR inserting the pos-th size into the descriptor155void MemRefDescriptor::setSize(OpBuilder &builder, Location loc, unsigned pos,156                               Value size) {157  value = LLVM::InsertValueOp::create(158      builder, loc, value, size,159      ArrayRef<int64_t>({kSizePosInMemRefDescriptor, pos}));160}161 162void MemRefDescriptor::setConstantSize(OpBuilder &builder, Location loc,163                                       unsigned pos, uint64_t size) {164  setSize(builder, loc, pos,165          createIndexAttrConstant(builder, loc, indexType, size));166}167 168/// Builds IR extracting the pos-th stride from the descriptor.169Value MemRefDescriptor::stride(OpBuilder &builder, Location loc, unsigned pos) {170  return LLVM::ExtractValueOp::create(171      builder, loc, value,172      ArrayRef<int64_t>({kStridePosInMemRefDescriptor, pos}));173}174 175/// Builds IR inserting the pos-th stride into the descriptor176void MemRefDescriptor::setStride(OpBuilder &builder, Location loc, unsigned pos,177                                 Value stride) {178  value = LLVM::InsertValueOp::create(179      builder, loc, value, stride,180      ArrayRef<int64_t>({kStridePosInMemRefDescriptor, pos}));181}182 183void MemRefDescriptor::setConstantStride(OpBuilder &builder, Location loc,184                                         unsigned pos, uint64_t stride) {185  setStride(builder, loc, pos,186            createIndexAttrConstant(builder, loc, indexType, stride));187}188 189LLVM::LLVMPointerType MemRefDescriptor::getElementPtrType() {190  return cast<LLVM::LLVMPointerType>(191      cast<LLVM::LLVMStructType>(value.getType())192          .getBody()[kAlignedPtrPosInMemRefDescriptor]);193}194 195Value MemRefDescriptor::bufferPtr(OpBuilder &builder, Location loc,196                                  const LLVMTypeConverter &converter,197                                  MemRefType type) {198  // When we convert to LLVM, the input memref must have been normalized199  // beforehand. Hence, this call is guaranteed to work.200  auto [strides, offsetCst] = type.getStridesAndOffset();201 202  Value ptr = alignedPtr(builder, loc);203  // For zero offsets, we already have the base pointer.204  if (offsetCst == 0)205    return ptr;206 207  // Otherwise add the offset to the aligned base.208  Type indexType = converter.getIndexType();209  Value offsetVal =210      ShapedType::isDynamic(offsetCst)211          ? offset(builder, loc)212          : createIndexAttrConstant(builder, loc, indexType, offsetCst);213  Type elementType = converter.convertType(type.getElementType());214  ptr = LLVM::GEPOp::create(builder, loc, ptr.getType(), elementType, ptr,215                            offsetVal);216  return ptr;217}218 219/// Creates a MemRef descriptor structure from a list of individual values220/// composing that descriptor, in the following order:221/// - allocated pointer;222/// - aligned pointer;223/// - offset;224/// - <rank> sizes;225/// - <rank> strides;226/// where <rank> is the MemRef rank as provided in `type`.227Value MemRefDescriptor::pack(OpBuilder &builder, Location loc,228                             const LLVMTypeConverter &converter,229                             MemRefType type, ValueRange values) {230  Type llvmType = converter.convertType(type);231  auto d = MemRefDescriptor::poison(builder, loc, llvmType);232 233  d.setAllocatedPtr(builder, loc, values[kAllocatedPtrPosInMemRefDescriptor]);234  d.setAlignedPtr(builder, loc, values[kAlignedPtrPosInMemRefDescriptor]);235  d.setOffset(builder, loc, values[kOffsetPosInMemRefDescriptor]);236 237  int64_t rank = type.getRank();238  for (unsigned i = 0; i < rank; ++i) {239    d.setSize(builder, loc, i, values[kSizePosInMemRefDescriptor + i]);240    d.setStride(builder, loc, i, values[kSizePosInMemRefDescriptor + rank + i]);241  }242 243  return d;244}245 246/// Builds IR extracting individual elements of a MemRef descriptor structure247/// and returning them as `results` list.248void MemRefDescriptor::unpack(OpBuilder &builder, Location loc, Value packed,249                              MemRefType type,250                              SmallVectorImpl<Value> &results) {251  int64_t rank = type.getRank();252  results.reserve(results.size() + getNumUnpackedValues(type));253 254  MemRefDescriptor d(packed);255  results.push_back(d.allocatedPtr(builder, loc));256  results.push_back(d.alignedPtr(builder, loc));257  results.push_back(d.offset(builder, loc));258  for (int64_t i = 0; i < rank; ++i)259    results.push_back(d.size(builder, loc, i));260  for (int64_t i = 0; i < rank; ++i)261    results.push_back(d.stride(builder, loc, i));262}263 264/// Returns the number of non-aggregate values that would be produced by265/// `unpack`.266unsigned MemRefDescriptor::getNumUnpackedValues(MemRefType type) {267  // Two pointers, offset, <rank> sizes, <rank> strides.268  return 3 + 2 * type.getRank();269}270 271//===----------------------------------------------------------------------===//272// MemRefDescriptorView implementation.273//===----------------------------------------------------------------------===//274 275MemRefDescriptorView::MemRefDescriptorView(ValueRange range)276    : rank((range.size() - kSizePosInMemRefDescriptor) / 2), elements(range) {}277 278Value MemRefDescriptorView::allocatedPtr() {279  return elements[kAllocatedPtrPosInMemRefDescriptor];280}281 282Value MemRefDescriptorView::alignedPtr() {283  return elements[kAlignedPtrPosInMemRefDescriptor];284}285 286Value MemRefDescriptorView::offset() {287  return elements[kOffsetPosInMemRefDescriptor];288}289 290Value MemRefDescriptorView::size(unsigned pos) {291  return elements[kSizePosInMemRefDescriptor + pos];292}293 294Value MemRefDescriptorView::stride(unsigned pos) {295  return elements[kSizePosInMemRefDescriptor + rank + pos];296}297 298//===----------------------------------------------------------------------===//299// UnrankedMemRefDescriptor implementation300//===----------------------------------------------------------------------===//301 302/// Construct a helper for the given descriptor value.303UnrankedMemRefDescriptor::UnrankedMemRefDescriptor(Value descriptor)304    : StructBuilder(descriptor) {}305 306/// Builds IR creating an `undef` value of the descriptor type.307UnrankedMemRefDescriptor UnrankedMemRefDescriptor::poison(OpBuilder &builder,308                                                          Location loc,309                                                          Type descriptorType) {310  Value descriptor = LLVM::PoisonOp::create(builder, loc, descriptorType);311  return UnrankedMemRefDescriptor(descriptor);312}313Value UnrankedMemRefDescriptor::rank(OpBuilder &builder, Location loc) const {314  return extractPtr(builder, loc, kRankInUnrankedMemRefDescriptor);315}316void UnrankedMemRefDescriptor::setRank(OpBuilder &builder, Location loc,317                                       Value v) {318  setPtr(builder, loc, kRankInUnrankedMemRefDescriptor, v);319}320Value UnrankedMemRefDescriptor::memRefDescPtr(OpBuilder &builder,321                                              Location loc) const {322  return extractPtr(builder, loc, kPtrInUnrankedMemRefDescriptor);323}324void UnrankedMemRefDescriptor::setMemRefDescPtr(OpBuilder &builder,325                                                Location loc, Value v) {326  setPtr(builder, loc, kPtrInUnrankedMemRefDescriptor, v);327}328 329/// Builds IR populating an unranked MemRef descriptor structure from a list330/// of individual constituent values in the following order:331/// - rank of the memref;332/// - pointer to the memref descriptor.333Value UnrankedMemRefDescriptor::pack(OpBuilder &builder, Location loc,334                                     const LLVMTypeConverter &converter,335                                     UnrankedMemRefType type,336                                     ValueRange values) {337  Type llvmType = converter.convertType(type);338  auto d = UnrankedMemRefDescriptor::poison(builder, loc, llvmType);339 340  d.setRank(builder, loc, values[kRankInUnrankedMemRefDescriptor]);341  d.setMemRefDescPtr(builder, loc, values[kPtrInUnrankedMemRefDescriptor]);342  return d;343}344 345/// Builds IR extracting individual elements that compose an unranked memref346/// descriptor and returns them as `results` list.347void UnrankedMemRefDescriptor::unpack(OpBuilder &builder, Location loc,348                                      Value packed,349                                      SmallVectorImpl<Value> &results) {350  UnrankedMemRefDescriptor d(packed);351  results.reserve(results.size() + 2);352  results.push_back(d.rank(builder, loc));353  results.push_back(d.memRefDescPtr(builder, loc));354}355 356Value UnrankedMemRefDescriptor::computeSize(357    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,358    UnrankedMemRefDescriptor desc, unsigned addressSpace) {359  // Cache the index type.360  Type indexType = typeConverter.getIndexType();361 362  // Initialize shared constants.363  Value one = createIndexAttrConstant(builder, loc, indexType, 1);364  Value two = createIndexAttrConstant(builder, loc, indexType, 2);365  Value indexSize = createIndexAttrConstant(366      builder, loc, indexType,367      llvm::divideCeil(typeConverter.getIndexTypeBitwidth(), 8));368 369  // Emit IR computing the memory necessary to store the descriptor. This370  // assumes the descriptor to be371  //   { type*, type*, index, index[rank], index[rank] }372  // and densely packed, so the total size is373  //   2 * sizeof(pointer) + (1 + 2 * rank) * sizeof(index).374  // TODO: consider including the actual size (including eventual padding due375  // to data layout) into the unranked descriptor.376  Value pointerSize = createIndexAttrConstant(377      builder, loc, indexType,378      llvm::divideCeil(typeConverter.getPointerBitwidth(addressSpace), 8));379  Value doublePointerSize =380      LLVM::MulOp::create(builder, loc, indexType, two, pointerSize);381 382  // (1 + 2 * rank) * sizeof(index)383  Value rank = desc.rank(builder, loc);384  Value doubleRank = LLVM::MulOp::create(builder, loc, indexType, two, rank);385  Value doubleRankIncremented =386      LLVM::AddOp::create(builder, loc, indexType, doubleRank, one);387  Value rankIndexSize = LLVM::MulOp::create(builder, loc, indexType,388                                            doubleRankIncremented, indexSize);389 390  // Total allocation size.391  Value allocationSize = LLVM::AddOp::create(builder, loc, indexType,392                                             doublePointerSize, rankIndexSize);393  return allocationSize;394}395 396Value UnrankedMemRefDescriptor::allocatedPtr(397    OpBuilder &builder, Location loc, Value memRefDescPtr,398    LLVM::LLVMPointerType elemPtrType) {399  return LLVM::LoadOp::create(builder, loc, elemPtrType, memRefDescPtr);400}401 402void UnrankedMemRefDescriptor::setAllocatedPtr(403    OpBuilder &builder, Location loc, Value memRefDescPtr,404    LLVM::LLVMPointerType elemPtrType, Value allocatedPtr) {405  LLVM::StoreOp::create(builder, loc, allocatedPtr, memRefDescPtr);406}407 408static std::pair<Value, Type>409castToElemPtrPtr(OpBuilder &builder, Location loc, Value memRefDescPtr,410                 LLVM::LLVMPointerType elemPtrType) {411  auto elemPtrPtrType = LLVM::LLVMPointerType::get(builder.getContext());412  return {memRefDescPtr, elemPtrPtrType};413}414 415Value UnrankedMemRefDescriptor::alignedPtr(416    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,417    Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType) {418  auto [elementPtrPtr, elemPtrPtrType] =419      castToElemPtrPtr(builder, loc, memRefDescPtr, elemPtrType);420 421  Value alignedGep =422      LLVM::GEPOp::create(builder, loc, elemPtrPtrType, elemPtrType,423                          elementPtrPtr, ArrayRef<LLVM::GEPArg>{1});424  return LLVM::LoadOp::create(builder, loc, elemPtrType, alignedGep);425}426 427void UnrankedMemRefDescriptor::setAlignedPtr(428    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,429    Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType, Value alignedPtr) {430  auto [elementPtrPtr, elemPtrPtrType] =431      castToElemPtrPtr(builder, loc, memRefDescPtr, elemPtrType);432 433  Value alignedGep =434      LLVM::GEPOp::create(builder, loc, elemPtrPtrType, elemPtrType,435                          elementPtrPtr, ArrayRef<LLVM::GEPArg>{1});436  LLVM::StoreOp::create(builder, loc, alignedPtr, alignedGep);437}438 439Value UnrankedMemRefDescriptor::offsetBasePtr(440    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,441    Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType) {442  auto [elementPtrPtr, elemPtrPtrType] =443      castToElemPtrPtr(builder, loc, memRefDescPtr, elemPtrType);444 445  return LLVM::GEPOp::create(builder, loc, elemPtrPtrType, elemPtrType,446                             elementPtrPtr, ArrayRef<LLVM::GEPArg>{2});447}448 449Value UnrankedMemRefDescriptor::offset(OpBuilder &builder, Location loc,450                                       const LLVMTypeConverter &typeConverter,451                                       Value memRefDescPtr,452                                       LLVM::LLVMPointerType elemPtrType) {453  Value offsetPtr =454      offsetBasePtr(builder, loc, typeConverter, memRefDescPtr, elemPtrType);455  return LLVM::LoadOp::create(builder, loc, typeConverter.getIndexType(),456                              offsetPtr);457}458 459void UnrankedMemRefDescriptor::setOffset(OpBuilder &builder, Location loc,460                                         const LLVMTypeConverter &typeConverter,461                                         Value memRefDescPtr,462                                         LLVM::LLVMPointerType elemPtrType,463                                         Value offset) {464  Value offsetPtr =465      offsetBasePtr(builder, loc, typeConverter, memRefDescPtr, elemPtrType);466  LLVM::StoreOp::create(builder, loc, offset, offsetPtr);467}468 469Value UnrankedMemRefDescriptor::sizeBasePtr(470    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,471    Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType) {472  Type indexTy = typeConverter.getIndexType();473  Type structTy = LLVM::LLVMStructType::getLiteral(474      indexTy.getContext(), {elemPtrType, elemPtrType, indexTy, indexTy});475  auto resultType = LLVM::LLVMPointerType::get(builder.getContext());476  return LLVM::GEPOp::create(builder, loc, resultType, structTy, memRefDescPtr,477                             ArrayRef<LLVM::GEPArg>{0, 3});478}479 480Value UnrankedMemRefDescriptor::size(OpBuilder &builder, Location loc,481                                     const LLVMTypeConverter &typeConverter,482                                     Value sizeBasePtr, Value index) {483 484  Type indexTy = typeConverter.getIndexType();485  auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());486 487  Value sizeStoreGep =488      LLVM::GEPOp::create(builder, loc, ptrType, indexTy, sizeBasePtr, index);489  return LLVM::LoadOp::create(builder, loc, indexTy, sizeStoreGep);490}491 492void UnrankedMemRefDescriptor::setSize(OpBuilder &builder, Location loc,493                                       const LLVMTypeConverter &typeConverter,494                                       Value sizeBasePtr, Value index,495                                       Value size) {496  Type indexTy = typeConverter.getIndexType();497  auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());498 499  Value sizeStoreGep =500      LLVM::GEPOp::create(builder, loc, ptrType, indexTy, sizeBasePtr, index);501  LLVM::StoreOp::create(builder, loc, size, sizeStoreGep);502}503 504Value UnrankedMemRefDescriptor::strideBasePtr(505    OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter,506    Value sizeBasePtr, Value rank) {507  Type indexTy = typeConverter.getIndexType();508  auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());509 510  return LLVM::GEPOp::create(builder, loc, ptrType, indexTy, sizeBasePtr, rank);511}512 513Value UnrankedMemRefDescriptor::stride(OpBuilder &builder, Location loc,514                                       const LLVMTypeConverter &typeConverter,515                                       Value strideBasePtr, Value index,516                                       Value stride) {517  Type indexTy = typeConverter.getIndexType();518  auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());519 520  Value strideStoreGep =521      LLVM::GEPOp::create(builder, loc, ptrType, indexTy, strideBasePtr, index);522  return LLVM::LoadOp::create(builder, loc, indexTy, strideStoreGep);523}524 525void UnrankedMemRefDescriptor::setStride(OpBuilder &builder, Location loc,526                                         const LLVMTypeConverter &typeConverter,527                                         Value strideBasePtr, Value index,528                                         Value stride) {529  Type indexTy = typeConverter.getIndexType();530  auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());531 532  Value strideStoreGep =533      LLVM::GEPOp::create(builder, loc, ptrType, indexTy, strideBasePtr, index);534  LLVM::StoreOp::create(builder, loc, stride, strideStoreGep);535}536