brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.8 KiB · 7e9318a Raw
100 lines · cpp
1//===-- XeVMToLLVMIRTranslation.cpp - Translate XeVM to LLVM IR -*- C++ -*-===//2//3// This file is licensed 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 a translation between the MLIR XeVM dialect and10// LLVM IR.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Target/LLVMIR/Dialect/XeVM/XeVMToLLVMIRTranslation.h"15#include "mlir/Dialect/LLVMIR/XeVMDialect.h"16#include "mlir/IR/BuiltinAttributes.h"17#include "mlir/IR/Operation.h"18#include "mlir/Target/LLVMIR/ModuleTranslation.h"19 20#include "llvm/ADT/TypeSwitch.h"21#include "llvm/IR/Constants.h"22#include "llvm/IR/LLVMContext.h"23#include "llvm/IR/Metadata.h"24 25#include "llvm/IR/ConstantRange.h"26#include "llvm/IR/IRBuilder.h"27#include "llvm/Support/raw_ostream.h"28 29using namespace mlir;30using namespace mlir::LLVM;31 32namespace {33/// Implementation of the dialect interface that converts operations belonging34/// to the XeVM dialect to LLVM IR.35class XeVMDialectLLVMIRTranslationInterface36    : public LLVMTranslationDialectInterface {37public:38  using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;39 40  /// Attaches module-level metadata for functions marked as kernels.41  LogicalResult42  amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,43                 NamedAttribute attribute,44                 LLVM::ModuleTranslation &moduleTranslation) const final {45    StringRef attrName = attribute.getName().getValue();46    if (attrName == mlir::xevm::XeVMDialect::getCacheControlsAttrName()) {47      auto cacheControlsArray = dyn_cast<ArrayAttr>(attribute.getValue());48      if (cacheControlsArray.size() != 2) {49        return op->emitOpError(50            "Expected both L1 and L3 cache control attributes!");51      }52      if (instructions.size() != 1) {53        return op->emitOpError("Expecting a single instruction");54      }55      return handleDecorationCacheControl(instructions.front(),56                                          cacheControlsArray.getValue());57    }58    return success();59  }60 61private:62  static LogicalResult handleDecorationCacheControl(llvm::Instruction *inst,63                                                    ArrayRef<Attribute> attrs) {64    SmallVector<llvm::Metadata *> decorations;65    llvm::LLVMContext &ctx = inst->getContext();66    llvm::Type *i32Ty = llvm::IntegerType::getInt32Ty(ctx);67    llvm::transform(68        attrs, std::back_inserter(decorations),69        [&ctx, i32Ty](Attribute attr) -> llvm::Metadata * {70          auto valuesArray = dyn_cast<ArrayAttr>(attr).getValue();71          std::array<llvm::Metadata *, 4> metadata;72          llvm::transform(73              valuesArray, metadata.begin(), [i32Ty](Attribute valueAttr) {74                return llvm::ConstantAsMetadata::get(llvm::ConstantInt::get(75                    i32Ty, cast<IntegerAttr>(valueAttr).getValue()));76              });77          return llvm::MDNode::get(ctx, metadata);78        });79    constexpr llvm::StringLiteral decorationCacheControlMDName =80        "spirv.DecorationCacheControlINTEL";81    inst->setMetadata(decorationCacheControlMDName,82                      llvm::MDNode::get(ctx, decorations));83    return success();84  }85};86} // namespace87 88void mlir::registerXeVMDialectTranslation(::mlir::DialectRegistry &registry) {89  registry.insert<xevm::XeVMDialect>();90  registry.addExtension(+[](MLIRContext *ctx, xevm::XeVMDialect *dialect) {91    dialect->addInterfaces<XeVMDialectLLVMIRTranslationInterface>();92  });93}94 95void mlir::registerXeVMDialectTranslation(::mlir::MLIRContext &context) {96  DialectRegistry registry;97  registerXeVMDialectTranslation(registry);98  context.appendDialectRegistry(registry);99}100