MLIR 24.0.0git
OpenACCRuntimeUtils.cpp
Go to the documentation of this file.
1//===- OpenACCRuntimeUtils.cpp - OpenACC runtime call utilities -*- C++ -*-===//
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
10
11#include "mlir/IR/SymbolTable.h"
12#include "llvm/Support/ErrorHandling.h"
13
14#include <cassert>
15#include <optional>
16
17using namespace mlir;
18using namespace mlir::acc;
19
24
26 switch (fn) {
27#define ACC_RTL(Enum, Str, ...) \
28 case RuntimeFunction::Enum: \
29 return Str;
30#include "mlir/Dialect/OpenACC/OpenACCRuntimeFunctions.def"
31 }
32 llvm_unreachable("unknown ACC runtime function");
33}
34
35LLVM::LLVMFunctionType acc::getRuntimeFunctionType(MLIRContext *ctx,
36 RuntimeFunction fn) {
37 Type Void = LLVM::LLVMVoidType::get(ctx);
38 Type Ptr = LLVM::LLVMPointerType::get(ctx);
39 Type Int32 = IntegerType::get(ctx, 32);
40 Type Int64 = IntegerType::get(ctx, 64);
41
42 switch (fn) {
43#define ACC_RTL(Enum, Str, IsVarArg, ReturnType, ...) \
44 case RuntimeFunction::Enum: \
45 return LLVM::LLVMFunctionType::get(ReturnType, \
46 ArrayRef<Type>{__VA_ARGS__}, IsVarArg);
47#include "mlir/Dialect/OpenACC/OpenACCRuntimeFunctions.def"
48 }
49 llvm_unreachable("unknown ACC runtime function");
50}
51
53 switch (desc) {
54#define ACC_DESC_BEGIN(Enum, NameStr, DescKind) \
55 case DataDescriptor::Enum: \
56 return NameStr;
57#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
58 }
59 llvm_unreachable("unknown ACC data descriptor");
60}
61
63 switch (desc) {
64#define ACC_DESC_BEGIN(Enum, NameStr, DescKind) \
65 case DataDescriptor::Enum: \
66 return DescKind;
67#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
68 }
69 llvm_unreachable("unknown ACC data descriptor");
70}
71
72LLVM::LLVMStructType acc::getDataDescriptorType(MLIRContext *ctx,
73 DataDescriptor desc,
74 Type baseType) {
75 Type Int8 = IntegerType::get(ctx, 8);
76 Type Int32 = IntegerType::get(ctx, 32);
77 Type Int64 = IntegerType::get(ctx, 64);
78 Type Ptr = LLVM::LLVMPointerType::get(ctx);
79 Type Base = baseType;
80
81 switch (desc) {
82#define ACC_DESC_BEGIN(Enum, NameStr, DescKind) \
83 case DataDescriptor::Enum: { \
84 SmallVector<Type> fields;
85#define ACC_DESC_FIELD(Enum, Field, TypeToken) \
86 assert(TypeToken && "no type for descriptor field " #Field); \
87 fields.push_back(TypeToken);
88#define ACC_DESC_END(Enum) \
89 return LLVM::LLVMStructType::getLiteral(ctx, fields); \
90 }
91#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
92 }
93 llvm_unreachable("unknown ACC data descriptor");
94}
95
97 overrides[fn] = name.str();
98}
99
101 if (auto it = overrides.find(fn); it != overrides.end())
102 return it->second;
103 return getRuntimeFunctionName(fn);
104}
105
107 functionDisplayNameFn = std::move(fn);
108}
109
110std::string
111ACCRuntimeCallConfig::getFunctionDisplayName(StringRef mangledOrSymbol) const {
112 if (functionDisplayNameFn)
113 return functionDisplayNameFn(mangledOrSymbol);
114 return mangledOrSymbol.str();
115}
116
118 int64_t runtimeValue) {
119 deviceTypeRuntimeValues[type] = runtimeValue;
120}
121
123 if (auto it = deviceTypeRuntimeValues.find(type);
124 it != deviceTypeRuntimeValues.end())
125 return it->second;
126 llvm::report_fatal_error(
127 llvm::Twine("missing OpenACC runtime device-type mapping for ") +
128 stringifyDeviceType(type));
129}
130
132 int64_t runtimeValue) {
133 mapFlagRuntimeValues[flag] = runtimeValue;
134}
135
137 int64_t runtimeValue = 0;
138 for (unsigned bit = 0; bit != 32; ++bit) {
139 auto flag = static_cast<MapFlags>(1u << bit);
140 if (!bitEnumContainsAny(flags, flag))
141 continue;
142 auto it = mapFlagRuntimeValues.find(flag);
143 if (it == mapFlagRuntimeValues.end())
144 llvm::report_fatal_error(
145 llvm::Twine("missing OpenACC runtime map-flag mapping for ") +
146 stringifyMapFlags(flag));
147 runtimeValue |= it->second;
148 }
149 return runtimeValue;
150}
151
153 mapFlagsPostProcessFn = std::move(fn);
154}
155
157 MapFlags flags) const {
158 if (mapFlagsPostProcessFn)
159 return mapFlagsPostProcessFn(mapOp, flags);
160 return flags;
161}
162
163std::string ACCRuntimeCallConfig::formatMapFlags(MapFlags flags) const {
164 int64_t runtimeValue = getMapFlagsRuntimeValue(flags);
165 std::string rendered = stringifyMapFlags(flags);
166 rendered += " (";
167 rendered += std::to_string(runtimeValue);
168 rendered += " / 0x";
169 rendered += llvm::Twine::utohexstr(runtimeValue).str();
170 rendered += ")";
171 return rendered;
172}
173
175 asyncSyncRuntimeValue = runtimeValue;
176}
177
179 return asyncSyncRuntimeValue;
180}
181
183 asyncNoValueRuntimeValue = runtimeValue;
184}
185
187 return asyncNoValueRuntimeValue;
188}
189
192 declareBinaryDescriptorFn = std::move(fn);
193}
194
196 Location loc, OpBuilder &builder) const {
197 if (declareBinaryDescriptorFn)
198 return declareBinaryDescriptorFn(loc, builder);
199 return LLVM::ZeroOp::create(builder, loc,
200 LLVM::LLVMPointerType::get(builder.getContext()));
201}
202
204 ACCRuntimeCallConfig &config) {
205 for (uint32_t value = 0; value <= getMaxEnumValForDeviceType(); ++value)
206 if (std::optional<DeviceType> type = symbolizeDeviceType(value))
207 config.setDeviceTypeRuntimeValue(*type, value);
208}
209
211 for (unsigned bit = 0; bit != 32; ++bit) {
212 uint32_t value = 1u << bit;
213 if (std::optional<MapFlags> flag = symbolizeMapFlags(value))
214 config.setMapFlagRuntimeValue(*flag, value);
215 }
216}
217
218FailureOr<LLVM::CallOp>
220 Region &globalSymbolRegion, SymbolTable &symbolTable,
221 RuntimeFunction fn, const ACCRuntimeCallConfig &config,
222 ArrayRef<Value> arguments) {
223 MLIRContext *ctx = builder.getContext();
224 LLVM::LLVMFunctionType fnTy = getRuntimeFunctionType(ctx, fn);
225 StringRef symbolName = config.getName(fn);
226
227 auto func = symbolTable.lookup<LLVM::LLVMFuncOp>(symbolName);
228 if (func) {
229 // An existing declaration with a different signature cannot be called with
230 // the arguments expected by the runtime entry point.
231 if (func.getFunctionType() != fnTy)
232 return emitError(loc) << "OpenACC runtime function '" << symbolName
233 << "' is already declared with signature "
234 << func.getFunctionType() << ", expected " << fnTy;
235 } else {
236 OpBuilder moduleBuilder =
237 OpBuilder::atBlockEnd(&globalSymbolRegion.front());
238 func = LLVM::LLVMFuncOp::create(moduleBuilder, loc, symbolName, fnTy);
239 symbolTable.insert(func);
240 }
241
242 return LLVM::CallOp::create(builder, loc, func, arguments);
243}
MLIRContext * getContext() const
Definition Builders.h:56
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
static OpBuilder atBlockEnd(Block *block, Listener *listener=nullptr)
Create a builder and set the insertion point to after the last operation in the block but still insid...
Definition Builders.h:249
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & front()
Definition Region.h:65
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
Definition SymbolTable.h:25
Operation * lookup(StringRef name) const
Look up a symbol with the specified name, returning null if no such name exists.
StringAttr insert(Operation *symbol, Block::iterator insertPt={})
Insert a new symbol into the table, and rename it as necessary to avoid collisions.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Configuration for OpenACC to LLVM runtime lowering.
StringRef getName(RuntimeFunction fn) const
std::string formatMapFlags(MapFlags flags) const
Renders flags for a diagnostic as the names of the set bits and the decimal and hexadecimal encoding ...
std::string getFunctionDisplayName(StringRef mangledOrSymbol) const
void setName(RuntimeFunction fn, StringRef name)
Value createDeclareBinaryDescriptor(Location loc, OpBuilder &builder) const
void setAsyncSyncRuntimeValue(int64_t runtimeValue)
Runtime encoding of acc_async_sync, used when an operation carries no async clause.
void setDeviceTypeRuntimeValue(DeviceType type, int64_t runtimeValue)
Map an OpenACC dialect DeviceType to the integer encoding expected by the target runtime.
void setMapFlagRuntimeValue(MapFlags flag, int64_t runtimeValue)
Map a single OpenACC dialect MapFlags bit to the bit the target runtime gives the same meaning.
int64_t getDeviceTypeRuntimeValue(DeviceType type) const
std::function< MapFlags(Operation *mapOp, MapFlags flags)> MapFlagsPostProcessFn
Adjust the flags of a mapping before they are encoded.
void setFunctionDisplayNameFn(FunctionDisplayNameFn fn)
std::function< std::string(StringRef)> FunctionDisplayNameFn
int64_t getMapFlagsRuntimeValue(MapFlags flags) const
std::function< Value(Location, OpBuilder &)> DeclareBinaryDescriptorFn
Materialize the target-specific binary descriptor passed to __tgt_acc_declare.
MapFlags postProcessMapFlags(Operation *mapOp, MapFlags flags) const
void setMapFlagsPostProcessFn(MapFlagsPostProcessFn fn)
void setDeclareBinaryDescriptorFn(DeclareBinaryDescriptorFn fn)
void setAsyncNoValueRuntimeValue(int64_t runtimeValue)
Runtime encoding of acc_async_noval, used for an async clause without an argument.
LLVM::LLVMFunctionType getRuntimeFunctionType(MLIRContext *ctx, RuntimeFunction fn)
Builds the LLVM function type for fn in ctx.
void populateDialectIdentityDeviceTypeMapping(ACCRuntimeCallConfig &config)
Install a device-type mapping that uses OpenACC dialect enum ordinals as the runtime encoding.
StringRef getRuntimeFunctionName(RuntimeFunction fn)
Returns the default runtime symbol name for fn.
LLVM::LLVMStructType getDataDescriptorType(MLIRContext *ctx, DataDescriptor desc, Type baseType={})
Builds the LLVM type of desc in ctx.
DataDescKind getDataDescriptorKind(DataDescriptor desc)
Returns the descriptor kind the runtime reads from the version field of desc.
StringRef getDataDescriptorName(DataDescriptor desc)
Returns the name of the runtime type desc materializes.
DataDescriptor
IDs for the argument descriptors of the OpenACC data entry points, declared in OpenACCRuntimeDescript...
RuntimeFunction
IDs for OpenACC compiler-to-runtime entry points (__tgt_acc_*).
void populateDialectIdentityMapFlagsMapping(ACCRuntimeCallConfig &config)
Install a map-flag mapping that uses the OpenACC dialect bit positions as the runtime encoding,...
FailureOr< LLVM::CallOp > createRuntimeCall(Location loc, OpBuilder &builder, Region &globalSymbolRegion, SymbolTable &symbolTable, RuntimeFunction fn, const ACCRuntimeCallConfig &config, ArrayRef< Value > arguments)
Declares (if needed) and returns a call to the runtime function identified by fn using the name from ...
Include the generated interface declarations.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.