blob: 42393218f870523cb8b9c2e65d86b80228dbfe9f [file] [edit]
//===- ROCDLToLLVMIRTranslation.cpp - Translate ROCDL to LLVM IR ----------===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
//
// This file implements a translation between the MLIR ROCDL dialect and
// LLVM IR.
//
//===----------------------------------------------------------------------===//
#include "mlir/Target/LLVMIR/Dialect/ROCDL/ROCDLToLLVMIRTranslation.h"
#include "mlir/Dialect/LLVMIR/ROCDLDialect.h"
#include "mlir/IR/BuiltinAttributes.h"
#include "mlir/IR/Operation.h"
#include "mlir/Target/LLVMIR/ModuleTranslation.h"
#include "llvm/IR/IRBuilder.h"
#include "llvm/IR/IntrinsicsAMDGPU.h"
#include "llvm/Support/raw_ostream.h"
#include <cstdint>
using namespace mlir;
using namespace mlir::LLVM;
using mlir::LLVM::detail::createIntrinsicCall;
namespace {
/// Implementation of the dialect interface that converts operations belonging
/// to the ROCDL dialect to LLVM IR.
class ROCDLDialectLLVMIRTranslationInterface
: public LLVMTranslationDialectInterface {
public:
using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
/// Translates the given operation to LLVM IR using the provided IR builder
/// and saving the state in `moduleTranslation`.
LogicalResult
convertOperation(Operation *op, llvm::IRBuilderBase &builder,
LLVM::ModuleTranslation &moduleTranslation) const final {
Operation &opInst = *op;
#include "mlir/Dialect/LLVMIR/ROCDLConversions.inc"
return failure();
}
/// Attaches module-level metadata for functions marked as kernels.
LogicalResult
amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
NamedAttribute attribute,
LLVM::ModuleTranslation &moduleTranslation) const final {
auto *dialect = dyn_cast<ROCDL::ROCDLDialect>(attribute.getNameDialect());
llvm::LLVMContext &llvmContext = moduleTranslation.getLLVMContext();
if (dialect->getKernelAttrHelper().getName() == attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
;
// For GPU kernels,
// 1. Insert AMDGPU_KERNEL calling convention.
// 2. Insert amdgpu-flat-work-group-size(1, 256) attribute unless the user
// has overriden this value - 256 is the default in clang
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvmFunc->setCallingConv(llvm::CallingConv::AMDGPU_KERNEL);
if (!llvmFunc->hasFnAttribute("amdgpu-flat-work-group-size")) {
llvmFunc->addFnAttr("amdgpu-flat-work-group-size", "1,256");
}
// MLIR's GPU kernel APIs all assume and produce uniformly-sized
// workgroups, so the lowering of the `rocdl.kernel` marker encodes this
// assumption. This assumption may be overridden by setting
// `rocdl.uniform_work_group_size` on a given function.
if (!llvmFunc->hasFnAttribute("uniform-work-group-size"))
llvmFunc->addFnAttr("uniform-work-group-size");
}
// Override flat-work-group-size
// TODO: update clients to rocdl.flat_work_group_size instead,
// then remove this half of the branch
if (dialect->getMaxFlatWorkGroupSizeAttrHelper().getName() ==
attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<IntegerAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be an integer");
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvm::SmallString<8> llvmAttrValue;
llvm::raw_svector_ostream attrValueStream(llvmAttrValue);
attrValueStream << "1," << value.getInt();
llvmFunc->addFnAttr("amdgpu-flat-work-group-size", llvmAttrValue);
}
if (dialect->getWavesPerEuAttrHelper().getName() == attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<IntegerAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be an integer");
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvm::SmallString<8> llvmAttrValue;
llvm::raw_svector_ostream attrValueStream(llvmAttrValue);
attrValueStream << value.getInt();
llvmFunc->addFnAttr("amdgpu-waves-per-eu", llvmAttrValue);
}
if (dialect->getFlatWorkGroupSizeAttrHelper().getName() ==
attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<StringAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be a string");
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvm::SmallString<8> llvmAttrValue;
llvmAttrValue.append(value.getValue());
llvmFunc->addFnAttr("amdgpu-flat-work-group-size", llvmAttrValue);
}
if (ROCDL::ROCDLDialect::getUniformWorkGroupSizeAttrName() ==
attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<BoolAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be a boolean");
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
if (value.getValue())
llvmFunc->addFnAttr("uniform-work-group-size");
else
llvmFunc->removeFnAttr("uniform-work-group-size");
}
if (dialect->getUnsafeFpAtomicsAttrHelper().getName() ==
attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<BoolAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be a boolean");
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvmFunc->addFnAttr("amdgpu-unsafe-fp-atomics",
value.getValue() ? "true" : "false");
}
// Set reqd_work_group_size metadata
if (dialect->getReqdWorkGroupSizeAttrHelper().getName() ==
attribute.getName()) {
auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
if (!func)
return op->emitOpError(Twine(attribute.getName()) +
" is only supported on `llvm.func` operations");
auto value = dyn_cast<DenseI32ArrayAttr>(attribute.getValue());
if (!value)
return op->emitOpError(Twine(attribute.getName()) +
" must be a dense i32 array attribute");
if (value.asArrayRef().size() != 3)
return op->emitOpError(Twine(attribute.getName()) +
" must contain exactly three values");
uint64_t FlatWorkGroupSize = 1;
SmallVector<llvm::Metadata *, 3> metadata;
llvm::Type *i32 = llvm::IntegerType::get(llvmContext, 32);
for (int32_t i : value.asArrayRef()) {
FlatWorkGroupSize *= static_cast<uint32_t>(i);
llvm::Constant *constant = llvm::ConstantInt::get(i32, i);
metadata.push_back(llvm::ConstantAsMetadata::get(constant));
}
llvm::Function *llvmFunc =
moduleTranslation.lookupFunction(func.getName());
llvm::SmallString<16> expectedFlatWorkGroupSize;
llvm::raw_svector_ostream attrValueStream(expectedFlatWorkGroupSize);
attrValueStream << FlatWorkGroupSize << "," << FlatWorkGroupSize;
StringRef flatAttrName =
dialect->getFlatWorkGroupSizeAttrHelper().getName();
if (auto flatAttr =
dyn_cast_if_present<StringAttr>(op->getAttr(flatAttrName))) {
if (flatAttr.getValue() != expectedFlatWorkGroupSize)
return op->emitOpError(Twine(flatAttrName) +
" must match rocdl.reqd_work_group_size");
}
StringRef maxFlatAttrName =
dialect->getMaxFlatWorkGroupSizeAttrHelper().getName();
if (auto maxFlatAttr =
dyn_cast_if_present<IntegerAttr>(op->getAttr(maxFlatAttrName))) {
llvm::SmallString<16> expectedMaxFlatWorkGroupSize;
llvm::raw_svector_ostream maxAttrValueStream(
expectedMaxFlatWorkGroupSize);
maxAttrValueStream << "1," << maxFlatAttr.getInt();
if (expectedMaxFlatWorkGroupSize != expectedFlatWorkGroupSize)
return op->emitOpError(Twine(maxFlatAttrName) +
" must match rocdl.reqd_work_group_size");
}
llvmFunc->addFnAttr("amdgpu-flat-work-group-size",
expectedFlatWorkGroupSize);
llvm::MDNode *node = llvm::MDNode::get(llvmContext, metadata);
llvmFunc->setMetadata("reqd_work_group_size", node);
}
// Atomic and nontemporal metadata
if (dialect->getLastUseAttrHelper().getName() == attribute.getName()) {
for (llvm::Instruction *i : instructions)
i->setMetadata("amdgpu.last.use", llvm::MDNode::get(llvmContext, {}));
}
if (dialect->getNoRemoteMemoryAttrHelper().getName() ==
attribute.getName()) {
for (llvm::Instruction *i : instructions)
i->setMetadata("amdgpu.no.remote.memory",
llvm::MDNode::get(llvmContext, {}));
}
if (dialect->getNoFineGrainedMemoryAttrHelper().getName() ==
attribute.getName()) {
for (llvm::Instruction *i : instructions)
i->setMetadata("amdgpu.no.fine.grained.memory",
llvm::MDNode::get(llvmContext, {}));
}
if (dialect->getIgnoreDenormalModeAttrHelper().getName() ==
attribute.getName()) {
for (llvm::Instruction *i : instructions)
i->setMetadata("amdgpu.ignore.denormal.mode",
llvm::MDNode::get(llvmContext, {}));
}
return success();
}
};
} // namespace
void mlir::registerROCDLDialectTranslation(DialectRegistry &registry) {
registry.insert<ROCDL::ROCDLDialect>();
registry.addExtension(+[](MLIRContext *ctx, ROCDL::ROCDLDialect *dialect) {
dialect->addInterfaces<ROCDLDialectLLVMIRTranslationInterface>();
});
}
void mlir::registerROCDLDialectTranslation(MLIRContext &context) {
DialectRegistry registry;
registerROCDLDialectTranslation(registry);
context.appendDialectRegistry(registry);
}