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