20#include "llvm/IR/IRBuilder.h"
21#include "llvm/IR/IntrinsicsAMDGPU.h"
22#include "llvm/IR/LLVMContext.h"
23#include "llvm/Support/raw_ostream.h"
33class ROCDLDialectLLVMIRTranslationInterface
34 :
public LLVMTranslationDialectInterface {
36 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
41 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
42 LLVM::ModuleTranslation &moduleTranslation)
const final {
43 Operation &opInst = *op;
44#include "mlir/Dialect/LLVMIR/ROCDLConversions.inc"
51 amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
52 NamedAttribute attribute,
53 LLVM::ModuleTranslation &moduleTranslation)
const final {
54 auto *dialect = dyn_cast<ROCDL::ROCDLDialect>(attribute.getNameDialect());
55 llvm::LLVMContext &llvmContext = moduleTranslation.getLLVMContext();
56 if (dialect->getKernelAttrHelper().getName() == attribute.getName()) {
57 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
59 return op->emitOpError(Twine(attribute.getName()) +
60 " is only supported on `llvm.func` operations");
67 llvm::Function *llvmFunc =
68 moduleTranslation.lookupFunction(func.getName());
69 llvmFunc->setCallingConv(llvm::CallingConv::AMDGPU_KERNEL);
70 if (!llvmFunc->hasFnAttribute(
"amdgpu-flat-work-group-size")) {
71 llvmFunc->addFnAttr(
"amdgpu-flat-work-group-size",
"1,256");
78 if (!llvmFunc->hasFnAttribute(
"uniform-work-group-size"))
79 llvmFunc->addFnAttr(
"uniform-work-group-size");
84 if (dialect->getMaxFlatWorkGroupSizeAttrHelper().getName() ==
85 attribute.getName()) {
86 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
88 return op->emitOpError(Twine(attribute.getName()) +
89 " is only supported on `llvm.func` operations");
90 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
92 return op->emitOpError(Twine(attribute.getName()) +
93 " must be an integer");
95 llvm::Function *llvmFunc =
96 moduleTranslation.lookupFunction(func.getName());
97 llvm::SmallString<8> llvmAttrValue;
98 llvm::raw_svector_ostream attrValueStream(llvmAttrValue);
99 attrValueStream <<
"1," << value.getInt();
100 llvmFunc->addFnAttr(
"amdgpu-flat-work-group-size", llvmAttrValue);
102 if (dialect->getWavesPerEuAttrHelper().getName() == attribute.getName()) {
103 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
105 return op->emitOpError(Twine(attribute.getName()) +
106 " is only supported on `llvm.func` operations");
107 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
109 return op->emitOpError(Twine(attribute.getName()) +
110 " must be an integer");
112 llvm::Function *llvmFunc =
113 moduleTranslation.lookupFunction(func.getName());
114 llvm::SmallString<8> llvmAttrValue;
115 llvm::raw_svector_ostream attrValueStream(llvmAttrValue);
116 attrValueStream << value.getInt();
117 llvmFunc->addFnAttr(
"amdgpu-waves-per-eu", llvmAttrValue);
119 if (dialect->getFlatWorkGroupSizeAttrHelper().getName() ==
120 attribute.getName()) {
121 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
123 return op->emitOpError(Twine(attribute.getName()) +
124 " is only supported on `llvm.func` operations");
125 auto value = dyn_cast<StringAttr>(attribute.getValue());
127 return op->emitOpError(Twine(attribute.getName()) +
128 " must be a string");
130 llvm::Function *llvmFunc =
131 moduleTranslation.lookupFunction(func.getName());
132 llvm::SmallString<8> llvmAttrValue;
133 llvmAttrValue.append(value.getValue());
134 llvmFunc->addFnAttr(
"amdgpu-flat-work-group-size", llvmAttrValue);
136 if (ROCDL::ROCDLDialect::getUniformWorkGroupSizeAttrName() ==
137 attribute.getName()) {
138 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
140 return op->emitOpError(Twine(attribute.getName()) +
141 " is only supported on `llvm.func` operations");
142 auto value = dyn_cast<BoolAttr>(attribute.getValue());
144 return op->emitOpError(Twine(attribute.getName()) +
145 " must be a boolean");
146 llvm::Function *llvmFunc =
147 moduleTranslation.lookupFunction(func.getName());
148 if (value.getValue())
149 llvmFunc->addFnAttr(
"uniform-work-group-size");
151 llvmFunc->removeFnAttr(
"uniform-work-group-size");
155 dialect->getXnackAttrHelper().getName() == attribute.getName();
157 dialect->getSrameccAttrHelper().getName() == attribute.getName();
158 if (isXnack || isSramecc) {
159 auto value = dyn_cast<BoolAttr>(attribute.getValue());
161 return op->emitOpError(Twine(attribute.getName()) +
162 " must be a boolean");
163 StringRef key = isXnack
164 ? ROCDL::ROCDLDialect::getModuleFlagKeyXnackName()
165 : ROCDL::ROCDLDialect::getModuleFlagKeySramEccName();
166 moduleTranslation.getLLVMModule()->addModuleFlag(
167 llvm::Module::Error, key,
168 llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvmContext),
171 if (dialect->getUnsafeFpAtomicsAttrHelper().getName() ==
172 attribute.getName()) {
173 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
175 return op->emitOpError(Twine(attribute.getName()) +
176 " is only supported on `llvm.func` operations");
177 auto value = dyn_cast<BoolAttr>(attribute.getValue());
179 return op->emitOpError(Twine(attribute.getName()) +
180 " must be a boolean");
181 llvm::Function *llvmFunc =
182 moduleTranslation.lookupFunction(func.getName());
183 llvmFunc->addFnAttr(
"amdgpu-unsafe-fp-atomics",
184 value.getValue() ?
"true" :
"false");
187 if (dialect->getReqdWorkGroupSizeAttrHelper().getName() ==
188 attribute.getName()) {
189 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
191 return op->emitOpError(Twine(attribute.getName()) +
192 " is only supported on `llvm.func` operations");
193 auto value = dyn_cast<DenseI32ArrayAttr>(attribute.getValue());
195 return op->emitOpError(Twine(attribute.getName()) +
196 " must be a dense i32 array attribute");
197 if (value.asArrayRef().size() != 3)
198 return op->emitOpError(Twine(attribute.getName()) +
199 " must contain exactly three values");
201 uint64_t FlatWorkGroupSize = 1;
202 SmallVector<llvm::Metadata *, 3> metadata;
203 llvm::Type *i32 = llvm::IntegerType::get(llvmContext, 32);
204 for (int32_t i : value.asArrayRef()) {
205 FlatWorkGroupSize *=
static_cast<uint32_t
>(i);
206 llvm::Constant *constant = llvm::ConstantInt::get(i32, i);
207 metadata.push_back(llvm::ConstantAsMetadata::get(constant));
209 llvm::Function *llvmFunc =
210 moduleTranslation.lookupFunction(func.getName());
211 llvm::SmallString<16> expectedFlatWorkGroupSize;
212 llvm::raw_svector_ostream attrValueStream(expectedFlatWorkGroupSize);
213 attrValueStream << FlatWorkGroupSize <<
"," << FlatWorkGroupSize;
215 StringRef flatAttrName =
216 dialect->getFlatWorkGroupSizeAttrHelper().getName();
217 if (
auto flatAttr = dyn_cast_if_present<StringAttr>(
218 op->getDiscardableAttr(flatAttrName))) {
219 if (flatAttr.getValue() != expectedFlatWorkGroupSize)
220 return op->emitOpError(Twine(flatAttrName) +
221 " must match rocdl.reqd_work_group_size");
224 StringRef maxFlatAttrName =
225 dialect->getMaxFlatWorkGroupSizeAttrHelper().getName();
226 if (
auto maxFlatAttr = dyn_cast_if_present<IntegerAttr>(
227 op->getDiscardableAttr(maxFlatAttrName))) {
228 llvm::SmallString<16> expectedMaxFlatWorkGroupSize;
229 llvm::raw_svector_ostream maxAttrValueStream(
230 expectedMaxFlatWorkGroupSize);
231 maxAttrValueStream <<
"1," << maxFlatAttr.getInt();
232 if (expectedMaxFlatWorkGroupSize != expectedFlatWorkGroupSize)
233 return op->emitOpError(Twine(maxFlatAttrName) +
234 " must match rocdl.reqd_work_group_size");
237 llvmFunc->addFnAttr(
"amdgpu-flat-work-group-size",
238 expectedFlatWorkGroupSize);
239 llvm::MDNode *node = llvm::MDNode::get(llvmContext, metadata);
240 llvmFunc->setMetadata(
"reqd_work_group_size", node);
244 if (dialect->getLastUseAttrHelper().getName() == attribute.getName()) {
245 for (llvm::Instruction *i : instructions)
246 i->setMetadata(
"amdgpu.last.use", llvm::MDNode::get(llvmContext, {}));
248 if (dialect->getNoRemoteMemoryAttrHelper().getName() ==
249 attribute.getName()) {
250 for (llvm::Instruction *i : instructions)
251 i->setMetadata(
"amdgpu.no.remote.memory",
252 llvm::MDNode::get(llvmContext, {}));
254 if (dialect->getNoFineGrainedMemoryAttrHelper().getName() ==
255 attribute.getName()) {
256 for (llvm::Instruction *i : instructions)
257 i->setMetadata(
"amdgpu.no.fine.grained.memory",
258 llvm::MDNode::get(llvmContext, {}));
260 if (dialect->getIgnoreDenormalModeAttrHelper().getName() ==
261 attribute.getName()) {
262 for (llvm::Instruction *i : instructions)
263 i->setMetadata(llvm::LLVMContext::MD_atomic_ignore_denormal_mode,
264 llvm::MDNode::get(llvmContext, {}));
273 registry.
insert<ROCDL::ROCDLDialect>();
275 dialect->addInterfaces<ROCDLDialectLLVMIRTranslationInterface>();
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
MLIRContext is the top-level object for a collection of MLIR operations.
void appendDialectRegistry(const DialectRegistry ®istry)
Append the contents of the given dialect registry to the registry associated with this context.
llvm::CallInst * createIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID intrinsic, ArrayRef< llvm::Value * > args={}, ArrayRef< llvm::Type * > tys={})
Creates a call to an LLVM IR intrinsic function with the given arguments.
Include the generated interface declarations.
void registerROCDLDialectTranslation(DialectRegistry ®istry)
Register the ROCDL dialect and the translation from it to the LLVM IR in the given registry;.