MLIR 24.0.0git
ROCDLToLLVMIRTranslation.cpp
Go to the documentation of this file.
1//===- ROCDLToLLVMIRTranslation.cpp - Translate ROCDL to LLVM IR ----------===//
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-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements a translation between the MLIR ROCDL dialect and
10// LLVM IR.
11//
12//===----------------------------------------------------------------------===//
13
17#include "mlir/IR/Operation.h"
19
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"
24#include <cstdint>
25
26using namespace mlir;
27using namespace mlir::LLVM;
29
30namespace {
31/// Implementation of the dialect interface that converts operations belonging
32/// to the ROCDL dialect to LLVM IR.
33class ROCDLDialectLLVMIRTranslationInterface
34 : public LLVMTranslationDialectInterface {
35public:
36 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
37
38 /// Translates the given operation to LLVM IR using the provided IR builder
39 /// and saving the state in `moduleTranslation`.
40 LogicalResult
41 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
42 LLVM::ModuleTranslation &moduleTranslation) const final {
43 Operation &opInst = *op;
44#include "mlir/Dialect/LLVMIR/ROCDLConversions.inc"
45
46 return failure();
47 }
48
49 /// Attaches module-level metadata for functions marked as kernels.
50 LogicalResult
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);
58 if (!func)
59 return op->emitOpError(Twine(attribute.getName()) +
60 " is only supported on `llvm.func` operations");
61 ;
62
63 // For GPU kernels,
64 // 1. Insert AMDGPU_KERNEL calling convention.
65 // 2. Insert amdgpu-flat-work-group-size(1, 256) attribute unless the user
66 // has overriden this value - 256 is the default in clang
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");
72 }
73
74 // MLIR's GPU kernel APIs all assume and produce uniformly-sized
75 // workgroups, so the lowering of the `rocdl.kernel` marker encodes this
76 // assumption. This assumption may be overridden by setting
77 // `rocdl.uniform_work_group_size` on a given function.
78 if (!llvmFunc->hasFnAttribute("uniform-work-group-size"))
79 llvmFunc->addFnAttr("uniform-work-group-size");
80 }
81 // Override flat-work-group-size
82 // TODO: update clients to rocdl.flat_work_group_size instead,
83 // then remove this half of the branch
84 if (dialect->getMaxFlatWorkGroupSizeAttrHelper().getName() ==
85 attribute.getName()) {
86 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
87 if (!func)
88 return op->emitOpError(Twine(attribute.getName()) +
89 " is only supported on `llvm.func` operations");
90 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
91 if (!value)
92 return op->emitOpError(Twine(attribute.getName()) +
93 " must be an integer");
94
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);
101 }
102 if (dialect->getWavesPerEuAttrHelper().getName() == attribute.getName()) {
103 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
104 if (!func)
105 return op->emitOpError(Twine(attribute.getName()) +
106 " is only supported on `llvm.func` operations");
107 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
108 if (!value)
109 return op->emitOpError(Twine(attribute.getName()) +
110 " must be an integer");
111
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);
118 }
119 if (dialect->getFlatWorkGroupSizeAttrHelper().getName() ==
120 attribute.getName()) {
121 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
122 if (!func)
123 return op->emitOpError(Twine(attribute.getName()) +
124 " is only supported on `llvm.func` operations");
125 auto value = dyn_cast<StringAttr>(attribute.getValue());
126 if (!value)
127 return op->emitOpError(Twine(attribute.getName()) +
128 " must be a string");
129
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);
135 }
136 if (ROCDL::ROCDLDialect::getUniformWorkGroupSizeAttrName() ==
137 attribute.getName()) {
138 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
139 if (!func)
140 return op->emitOpError(Twine(attribute.getName()) +
141 " is only supported on `llvm.func` operations");
142 auto value = dyn_cast<BoolAttr>(attribute.getValue());
143 if (!value)
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");
150 else
151 llvmFunc->removeFnAttr("uniform-work-group-size");
152 }
153
154 bool isXnack =
155 dialect->getXnackAttrHelper().getName() == attribute.getName();
156 bool isSramecc =
157 dialect->getSrameccAttrHelper().getName() == attribute.getName();
158 if (isXnack || isSramecc) {
159 auto value = dyn_cast<BoolAttr>(attribute.getValue());
160 if (!value)
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),
169 value.getValue()));
170 }
171 if (dialect->getUnsafeFpAtomicsAttrHelper().getName() ==
172 attribute.getName()) {
173 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
174 if (!func)
175 return op->emitOpError(Twine(attribute.getName()) +
176 " is only supported on `llvm.func` operations");
177 auto value = dyn_cast<BoolAttr>(attribute.getValue());
178 if (!value)
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");
185 }
186 // Set reqd_work_group_size metadata
187 if (dialect->getReqdWorkGroupSizeAttrHelper().getName() ==
188 attribute.getName()) {
189 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
190 if (!func)
191 return op->emitOpError(Twine(attribute.getName()) +
192 " is only supported on `llvm.func` operations");
193 auto value = dyn_cast<DenseI32ArrayAttr>(attribute.getValue());
194 if (!value)
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");
200
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));
208 }
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;
214
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");
222 }
223
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");
235 }
236
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);
241 }
242
243 // Atomic and nontemporal metadata
244 if (dialect->getLastUseAttrHelper().getName() == attribute.getName()) {
245 for (llvm::Instruction *i : instructions)
246 i->setMetadata("amdgpu.last.use", llvm::MDNode::get(llvmContext, {}));
247 }
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, {}));
253 }
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, {}));
259 }
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, {}));
265 }
266
267 return success();
268 }
269};
270} // namespace
271
273 registry.insert<ROCDL::ROCDLDialect>();
274 registry.addExtension(+[](MLIRContext *ctx, ROCDL::ROCDLDialect *dialect) {
275 dialect->addInterfaces<ROCDLDialectLLVMIRTranslationInterface>();
276 });
277}
278
return success()
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.
Definition MLIRContext.h:63
void appendDialectRegistry(const DialectRegistry &registry)
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 &registry)
Register the ROCDL dialect and the translation from it to the LLVM IR in the given registry;.