brintos

brintos / llvm-project-archived public Read only

0
0
Text · 21.0 KiB · f143a9e Raw
539 lines · cpp
1//===- GPUToLLVMSPV.cpp - Convert GPU operations to LLVM dialect ----------===//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/GPUToLLVMSPV/GPUToLLVMSPVPass.h"10 11#include "../GPUCommon/GPUOpsLowering.h"12#include "mlir/Conversion/GPUCommon/AttrToSPIRVConverter.h"13#include "mlir/Conversion/GPUCommon/GPUCommonPass.h"14#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"15#include "mlir/Conversion/LLVMCommon/LoweringOptions.h"16#include "mlir/Conversion/LLVMCommon/Pattern.h"17#include "mlir/Conversion/LLVMCommon/TypeConverter.h"18#include "mlir/Conversion/SPIRVCommon/AttrToLLVMConverter.h"19#include "mlir/Dialect/GPU/IR/GPUDialect.h"20#include "mlir/Dialect/LLVMIR/LLVMAttrs.h"21#include "mlir/Dialect/LLVMIR/LLVMDialect.h"22#include "mlir/Dialect/LLVMIR/LLVMTypes.h"23#include "mlir/Dialect/SPIRV/IR/SPIRVEnums.h"24#include "mlir/IR/BuiltinTypes.h"25#include "mlir/IR/Matchers.h"26#include "mlir/IR/PatternMatch.h"27#include "mlir/IR/SymbolTable.h"28#include "mlir/Pass/Pass.h"29#include "mlir/Support/LLVM.h"30#include "mlir/Transforms/DialectConversion.h"31 32#include "llvm/ADT/TypeSwitch.h"33#include "llvm/Support/FormatVariadic.h"34 35#define DEBUG_TYPE "gpu-to-llvm-spv"36 37using namespace mlir;38 39namespace mlir {40#define GEN_PASS_DEF_CONVERTGPUOPSTOLLVMSPVOPS41#include "mlir/Conversion/Passes.h.inc"42} // namespace mlir43 44//===----------------------------------------------------------------------===//45// Helper Functions46//===----------------------------------------------------------------------===//47 48static LLVM::LLVMFuncOp lookupOrCreateSPIRVFn(Operation *symbolTable,49                                              StringRef name,50                                              ArrayRef<Type> paramTypes,51                                              Type resultType, bool isMemNone,52                                              bool isConvergent) {53  auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(54      SymbolTable::lookupSymbolIn(symbolTable, name));55  if (!func) {56    OpBuilder b(symbolTable->getRegion(0));57    func = LLVM::LLVMFuncOp::create(58        b, symbolTable->getLoc(), name,59        LLVM::LLVMFunctionType::get(resultType, paramTypes));60    func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);61    func.setNoUnwind(true);62    func.setWillReturn(true);63 64    if (isMemNone) {65      // no externally observable effects66      constexpr auto noModRef = mlir::LLVM::ModRefInfo::NoModRef;67      auto memAttr = b.getAttr<LLVM::MemoryEffectsAttr>(68          /*other=*/noModRef,69          /*argMem=*/noModRef, /*inaccessibleMem=*/noModRef,70          /*errnoMem=*/noModRef,71          /*targetMem0=*/noModRef,72          /*targetMem1=*/noModRef);73      func.setMemoryEffectsAttr(memAttr);74    }75 76    func.setConvergent(isConvergent);77  }78  return func;79}80 81static LLVM::CallOp createSPIRVBuiltinCall(Location loc,82                                           ConversionPatternRewriter &rewriter,83                                           LLVM::LLVMFuncOp func,84                                           ValueRange args) {85  auto call = LLVM::CallOp::create(rewriter, loc, func, args);86  call.setCConv(func.getCConv());87  call.setConvergentAttr(func.getConvergentAttr());88  call.setNoUnwindAttr(func.getNoUnwindAttr());89  call.setWillReturnAttr(func.getWillReturnAttr());90  call.setMemoryEffectsAttr(func.getMemoryEffectsAttr());91  return call;92}93 94namespace {95//===----------------------------------------------------------------------===//96// Barriers97//===----------------------------------------------------------------------===//98 99/// Replace `gpu.barrier` with an `llvm.call` to `barrier` with100/// `CLK_LOCAL_MEM_FENCE` argument, indicating work-group memory scope:101/// ```102/// // gpu.barrier103/// %c1 = llvm.mlir.constant(1: i32) : i32104/// llvm.call spir_funccc @_Z7barrierj(%c1) : (i32) -> ()105/// ```106struct GPUBarrierConversion final : ConvertOpToLLVMPattern<gpu::BarrierOp> {107  using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;108 109  LogicalResult110  matchAndRewrite(gpu::BarrierOp op, OpAdaptor adaptor,111                  ConversionPatternRewriter &rewriter) const final {112    constexpr StringLiteral funcName = "_Z7barrierj";113 114    Operation *moduleOp = op->getParentWithTrait<OpTrait::SymbolTable>();115    assert(moduleOp && "Expecting module");116    Type flagTy = rewriter.getI32Type();117    Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();118    LLVM::LLVMFuncOp func =119        lookupOrCreateSPIRVFn(moduleOp, funcName, flagTy, voidTy,120                              /*isMemNone=*/false, /*isConvergent=*/true);121 122    // Value used by SPIR-V backend to represent `CLK_LOCAL_MEM_FENCE`.123    // See `llvm/lib/Target/SPIRV/SPIRVBuiltins.td`.124    constexpr int64_t localMemFenceFlag = 1;125    Location loc = op->getLoc();126    Value flag =127        LLVM::ConstantOp::create(rewriter, loc, flagTy, localMemFenceFlag);128    rewriter.replaceOp(op, createSPIRVBuiltinCall(loc, rewriter, func, flag));129    return success();130  }131};132 133//===----------------------------------------------------------------------===//134// SPIR-V Builtins135//===----------------------------------------------------------------------===//136 137/// Replace `gpu.*` with an `llvm.call` to the corresponding SPIR-V builtin with138/// a constant argument for the `dimension` attribute. Return type will depend139/// on index width option:140/// ```141/// // %thread_id_y = gpu.thread_id y142/// %c1 = llvm.mlir.constant(1: i32) : i32143/// %0 = llvm.call spir_funccc @_Z12get_local_idj(%c1) : (i32) -> i64144/// ```145struct LaunchConfigConversion : ConvertToLLVMPattern {146  LaunchConfigConversion(StringRef funcName, StringRef rootOpName,147                         MLIRContext *context,148                         const LLVMTypeConverter &typeConverter,149                         PatternBenefit benefit)150      : ConvertToLLVMPattern(rootOpName, context, typeConverter, benefit),151        funcName(funcName) {}152 153  virtual gpu::Dimension getDimension(Operation *op) const = 0;154 155  LogicalResult156  matchAndRewrite(Operation *op, ArrayRef<Value> operands,157                  ConversionPatternRewriter &rewriter) const final {158    Operation *moduleOp = op->getParentWithTrait<OpTrait::SymbolTable>();159    assert(moduleOp && "Expecting module");160    Type dimTy = rewriter.getI32Type();161    Type indexTy = getTypeConverter()->getIndexType();162    LLVM::LLVMFuncOp func = lookupOrCreateSPIRVFn(moduleOp, funcName, dimTy,163                                                  indexTy, /*isMemNone=*/true,164                                                  /*isConvergent=*/false);165 166    Location loc = op->getLoc();167    gpu::Dimension dim = getDimension(op);168    Value dimVal = LLVM::ConstantOp::create(rewriter, loc, dimTy,169                                            static_cast<int64_t>(dim));170    rewriter.replaceOp(op, createSPIRVBuiltinCall(loc, rewriter, func, dimVal));171    return success();172  }173 174  StringRef funcName;175};176 177template <typename SourceOp>178struct LaunchConfigOpConversion final : LaunchConfigConversion {179  static StringRef getFuncName();180 181  explicit LaunchConfigOpConversion(const LLVMTypeConverter &typeConverter,182                                    PatternBenefit benefit = 1)183      : LaunchConfigConversion(getFuncName(), SourceOp::getOperationName(),184                               &typeConverter.getContext(), typeConverter,185                               benefit) {}186 187  gpu::Dimension getDimension(Operation *op) const final {188    return cast<SourceOp>(op).getDimension();189  }190};191 192template <>193StringRef LaunchConfigOpConversion<gpu::BlockIdOp>::getFuncName() {194  return "_Z12get_group_idj";195}196 197template <>198StringRef LaunchConfigOpConversion<gpu::GridDimOp>::getFuncName() {199  return "_Z14get_num_groupsj";200}201 202template <>203StringRef LaunchConfigOpConversion<gpu::BlockDimOp>::getFuncName() {204  return "_Z14get_local_sizej";205}206 207template <>208StringRef LaunchConfigOpConversion<gpu::ThreadIdOp>::getFuncName() {209  return "_Z12get_local_idj";210}211 212template <>213StringRef LaunchConfigOpConversion<gpu::GlobalIdOp>::getFuncName() {214  return "_Z13get_global_idj";215}216 217//===----------------------------------------------------------------------===//218// Shuffles219//===----------------------------------------------------------------------===//220 221/// Replace `gpu.shuffle` with an `llvm.call` to the corresponding SPIR-V222/// builtin for `shuffleResult`, keeping `value` and `offset` arguments, and a223/// `true` constant for the `valid` result type. Conversion will only take place224/// if `width` is constant and equal to the `subgroup` pass option:225/// ```226/// // %0 = gpu.shuffle idx %value, %offset, %width : f64227/// %0 = llvm.call spir_funccc @_Z17sub_group_shuffledj(%value, %offset)228///     : (f64, i32) -> f64229/// ```230struct GPUShuffleConversion final : ConvertOpToLLVMPattern<gpu::ShuffleOp> {231  using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;232 233  static StringRef getBaseName(gpu::ShuffleMode mode) {234    switch (mode) {235    case gpu::ShuffleMode::IDX:236      return "sub_group_shuffle";237    case gpu::ShuffleMode::XOR:238      return "sub_group_shuffle_xor";239    case gpu::ShuffleMode::UP:240      return "sub_group_shuffle_up";241    case gpu::ShuffleMode::DOWN:242      return "sub_group_shuffle_down";243    }244    llvm_unreachable("Unhandled shuffle mode");245  }246 247  static std::optional<StringRef> getTypeMangling(Type type) {248    return TypeSwitch<Type, std::optional<StringRef>>(type)249        .Case<Float16Type>([](auto) { return "Dhj"; })250        .Case<Float32Type>([](auto) { return "fj"; })251        .Case<Float64Type>([](auto) { return "dj"; })252        .Case<IntegerType>([](auto intTy) -> std::optional<StringRef> {253          switch (intTy.getWidth()) {254          case 8:255            return "cj";256          case 16:257            return "sj";258          case 32:259            return "ij";260          case 64:261            return "lj";262          }263          return std::nullopt;264        })265        .Default(std::nullopt);266  }267 268  static std::optional<std::string> getFuncName(gpu::ShuffleMode mode,269                                                Type type) {270    StringRef baseName = getBaseName(mode);271    std::optional<StringRef> typeMangling = getTypeMangling(type);272    if (!typeMangling)273      return std::nullopt;274    return llvm::formatv("_Z{}{}{}", baseName.size(), baseName,275                         typeMangling.value());276  }277 278  /// Get the subgroup size from the target or return a default.279  static std::optional<int> getSubgroupSize(Operation *op) {280    auto parentFunc = op->getParentOfType<LLVM::LLVMFuncOp>();281    if (!parentFunc)282      return std::nullopt;283    return parentFunc.getIntelReqdSubGroupSize();284  }285 286  static bool hasValidWidth(gpu::ShuffleOp op) {287    llvm::APInt val;288    Value width = op.getWidth();289    return matchPattern(width, m_ConstantInt(&val)) &&290           val == getSubgroupSize(op);291  }292 293  static Value bitcastOrExtBeforeShuffle(Value oldVal, Location loc,294                                         ConversionPatternRewriter &rewriter) {295    return TypeSwitch<Type, Value>(oldVal.getType())296        .Case([&](BFloat16Type) {297          return LLVM::BitcastOp::create(rewriter, loc, rewriter.getI16Type(),298                                         oldVal);299        })300        .Case([&](IntegerType intTy) -> Value {301          if (intTy.getWidth() == 1)302            return LLVM::ZExtOp::create(rewriter, loc, rewriter.getI8Type(),303                                        oldVal);304          return oldVal;305        })306        .Default(oldVal);307  }308 309  static Value bitcastOrTruncAfterShuffle(Value oldVal, Type newTy,310                                          Location loc,311                                          ConversionPatternRewriter &rewriter) {312    return TypeSwitch<Type, Value>(newTy)313        .Case([&](BFloat16Type) {314          return LLVM::BitcastOp::create(rewriter, loc, newTy, oldVal);315        })316        .Case([&](IntegerType intTy) -> Value {317          if (intTy.getWidth() == 1)318            return LLVM::TruncOp::create(rewriter, loc, newTy, oldVal);319          return oldVal;320        })321        .Default(oldVal);322  }323 324  LogicalResult325  matchAndRewrite(gpu::ShuffleOp op, OpAdaptor adaptor,326                  ConversionPatternRewriter &rewriter) const final {327    if (!hasValidWidth(op))328      return rewriter.notifyMatchFailure(329          op, "shuffle width and subgroup size mismatch");330 331    Location loc = op->getLoc();332    Value inValue =333        bitcastOrExtBeforeShuffle(adaptor.getValue(), loc, rewriter);334    std::optional<std::string> funcName =335        getFuncName(op.getMode(), inValue.getType());336    if (!funcName)337      return rewriter.notifyMatchFailure(op, "unsupported value type");338 339    Operation *moduleOp = op->getParentWithTrait<OpTrait::SymbolTable>();340    assert(moduleOp && "Expecting module");341    Type valueType = inValue.getType();342    Type offsetType = adaptor.getOffset().getType();343    Type resultType = valueType;344    LLVM::LLVMFuncOp func = lookupOrCreateSPIRVFn(345        moduleOp, funcName.value(), {valueType, offsetType}, resultType,346        /*isMemNone=*/false, /*isConvergent=*/true);347 348    std::array<Value, 2> args{inValue, adaptor.getOffset()};349    Value result =350        createSPIRVBuiltinCall(loc, rewriter, func, args).getResult();351    Value resultOrConversion =352        bitcastOrTruncAfterShuffle(result, op.getType(0), loc, rewriter);353 354    Value trueVal =355        LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), true);356    rewriter.replaceOp(op, {resultOrConversion, trueVal});357    return success();358  }359};360 361class MemorySpaceToOpenCLMemorySpaceConverter final : public TypeConverter {362public:363  MemorySpaceToOpenCLMemorySpaceConverter(MLIRContext *ctx) {364    addConversion([](Type t) { return t; });365    addConversion([ctx](BaseMemRefType memRefType) -> std::optional<Type> {366      // Attach global addr space attribute to memrefs with no addr space attr367      Attribute memSpaceAttr = memRefType.getMemorySpace();368      if (memSpaceAttr)369        return std::nullopt;370 371      unsigned globalAddrspace = storageClassToAddressSpace(372          spirv::ClientAPI::OpenCL, spirv::StorageClass::CrossWorkgroup);373      Attribute addrSpaceAttr =374          IntegerAttr::get(IntegerType::get(ctx, 64), globalAddrspace);375      if (auto rankedType = dyn_cast<MemRefType>(memRefType)) {376        return MemRefType::get(memRefType.getShape(),377                               memRefType.getElementType(),378                               rankedType.getLayout(), addrSpaceAttr);379      }380      return UnrankedMemRefType::get(memRefType.getElementType(),381                                     addrSpaceAttr);382    });383    addConversion([this](FunctionType type) {384      auto inputs = llvm::map_to_vector(385          type.getInputs(), [this](Type ty) { return convertType(ty); });386      auto results = llvm::map_to_vector(387          type.getResults(), [this](Type ty) { return convertType(ty); });388      return FunctionType::get(type.getContext(), inputs, results);389    });390  }391};392 393//===----------------------------------------------------------------------===//394// Subgroup query ops.395//===----------------------------------------------------------------------===//396 397template <typename SubgroupOp>398struct GPUSubgroupOpConversion final : ConvertOpToLLVMPattern<SubgroupOp> {399  using ConvertOpToLLVMPattern<SubgroupOp>::ConvertOpToLLVMPattern;400  using ConvertToLLVMPattern::getTypeConverter;401 402  LogicalResult403  matchAndRewrite(SubgroupOp op, typename SubgroupOp::Adaptor adaptor,404                  ConversionPatternRewriter &rewriter) const final {405    constexpr StringRef funcName = [] {406      if constexpr (std::is_same_v<SubgroupOp, gpu::SubgroupIdOp>) {407        return "_Z16get_sub_group_id";408      } else if constexpr (std::is_same_v<SubgroupOp, gpu::LaneIdOp>) {409        return "_Z22get_sub_group_local_id";410      } else if constexpr (std::is_same_v<SubgroupOp, gpu::NumSubgroupsOp>) {411        return "_Z18get_num_sub_groups";412      } else if constexpr (std::is_same_v<SubgroupOp, gpu::SubgroupSizeOp>) {413        return "_Z18get_sub_group_size";414      }415    }();416 417    Operation *moduleOp =418        op->template getParentWithTrait<OpTrait::SymbolTable>();419    Type resultTy = rewriter.getI32Type();420    LLVM::LLVMFuncOp func =421        lookupOrCreateSPIRVFn(moduleOp, funcName, {}, resultTy,422                              /*isMemNone=*/false, /*isConvergent=*/false);423 424    Location loc = op->getLoc();425    Value result = createSPIRVBuiltinCall(loc, rewriter, func, {}).getResult();426 427    Type indexTy = getTypeConverter()->getIndexType();428    if (resultTy != indexTy) {429      if (indexTy.getIntOrFloatBitWidth() < resultTy.getIntOrFloatBitWidth()) {430        return failure();431      }432      result = LLVM::ZExtOp::create(rewriter, loc, indexTy, result);433    }434 435    rewriter.replaceOp(op, result);436    return success();437  }438};439 440//===----------------------------------------------------------------------===//441// GPU To LLVM-SPV Pass.442//===----------------------------------------------------------------------===//443 444struct GPUToLLVMSPVConversionPass final445    : impl::ConvertGpuOpsToLLVMSPVOpsBase<GPUToLLVMSPVConversionPass> {446  using Base::Base;447 448  void runOnOperation() final {449    MLIRContext *context = &getContext();450    RewritePatternSet patterns(context);451 452    LowerToLLVMOptions options(context);453    options.overrideIndexBitwidth(this->use64bitIndex ? 64 : 32);454    LLVMTypeConverter converter(context, options);455    LLVMConversionTarget target(*context);456 457    // Force OpenCL address spaces when they are not present458    {459      MemorySpaceToOpenCLMemorySpaceConverter converter(context);460      AttrTypeReplacer replacer;461      replacer.addReplacement([&converter](BaseMemRefType origType)462                                  -> std::optional<BaseMemRefType> {463        return converter.convertType<BaseMemRefType>(origType);464      });465 466      replacer.recursivelyReplaceElementsIn(getOperation(),467                                            /*replaceAttrs=*/true,468                                            /*replaceLocs=*/false,469                                            /*replaceTypes=*/true);470    }471 472    target.addIllegalOp<gpu::BarrierOp, gpu::BlockDimOp, gpu::BlockIdOp,473                        gpu::GPUFuncOp, gpu::GlobalIdOp, gpu::GridDimOp,474                        gpu::LaneIdOp, gpu::NumSubgroupsOp, gpu::ReturnOp,475                        gpu::ShuffleOp, gpu::SubgroupIdOp, gpu::SubgroupSizeOp,476                        gpu::ThreadIdOp, gpu::PrintfOp>();477 478    populateGpuToLLVMSPVConversionPatterns(converter, patterns);479    populateGpuMemorySpaceAttributeConversions(converter);480    patterns.add<GPUPrintfOpToLLVMCallLowering>(converter, /*addressSpace=*/2,481                                                LLVM::cconv::CConv::SPIR_FUNC,482                                                "_Z6printfPU3AS2Kcz");483 484    if (failed(applyPartialConversion(getOperation(), target,485                                      std::move(patterns))))486      signalPassFailure();487  }488};489} // namespace490 491//===----------------------------------------------------------------------===//492// GPU To LLVM-SPV Patterns.493//===----------------------------------------------------------------------===//494 495namespace mlir {496namespace {497static unsigned498gpuAddressSpaceToOCLAddressSpace(gpu::AddressSpace addressSpace) {499  constexpr spirv::ClientAPI clientAPI = spirv::ClientAPI::OpenCL;500  return storageClassToAddressSpace(clientAPI,501                                    addressSpaceToStorageClass(addressSpace));502}503} // namespace504 505void populateGpuToLLVMSPVConversionPatterns(506    const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {507  patterns.add<GPUBarrierConversion, GPUReturnOpLowering, GPUShuffleConversion,508               GPUSubgroupOpConversion<gpu::LaneIdOp>,509               GPUSubgroupOpConversion<gpu::NumSubgroupsOp>,510               GPUSubgroupOpConversion<gpu::SubgroupIdOp>,511               GPUSubgroupOpConversion<gpu::SubgroupSizeOp>,512               LaunchConfigOpConversion<gpu::BlockDimOp>,513               LaunchConfigOpConversion<gpu::BlockIdOp>,514               LaunchConfigOpConversion<gpu::GlobalIdOp>,515               LaunchConfigOpConversion<gpu::GridDimOp>,516               LaunchConfigOpConversion<gpu::ThreadIdOp>>(typeConverter);517  MLIRContext *context = &typeConverter.getContext();518  unsigned privateAddressSpace =519      gpuAddressSpaceToOCLAddressSpace(gpu::AddressSpace::Private);520  unsigned localAddressSpace =521      gpuAddressSpaceToOCLAddressSpace(gpu::AddressSpace::Workgroup);522  OperationName llvmFuncOpName(LLVM::LLVMFuncOp::getOperationName(), context);523  StringAttr kernelBlockSizeAttributeName =524      LLVM::LLVMFuncOp::getReqdWorkGroupSizeAttrName(llvmFuncOpName);525  patterns.add<GPUFuncOpLowering>(526      typeConverter,527      GPUFuncOpLoweringOptions{528          privateAddressSpace, localAddressSpace,529          /*kernelAttributeName=*/{}, kernelBlockSizeAttributeName,530          LLVM::CConv::SPIR_KERNEL, LLVM::CConv::SPIR_FUNC,531          /*encodeWorkgroupAttributionsAsArguments=*/true});532}533 534void populateGpuMemorySpaceAttributeConversions(TypeConverter &typeConverter) {535  populateGpuMemorySpaceAttributeConversions(typeConverter,536                                             gpuAddressSpaceToOCLAddressSpace);537}538} // namespace mlir539