12#include "llvm/Support/ErrorHandling.h"
27#define ACC_RTL(Enum, Str, ...) \
28 case RuntimeFunction::Enum: \
30#include "mlir/Dialect/OpenACC/OpenACCRuntimeFunctions.def"
32 llvm_unreachable(
"unknown ACC runtime function");
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);
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"
49 llvm_unreachable(
"unknown ACC runtime function");
54#define ACC_DESC_BEGIN(Enum, NameStr, DescKind) \
55 case DataDescriptor::Enum: \
57#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
59 llvm_unreachable(
"unknown ACC data descriptor");
64#define ACC_DESC_BEGIN(Enum, NameStr, DescKind) \
65 case DataDescriptor::Enum: \
67#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
69 llvm_unreachable(
"unknown ACC data descriptor");
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);
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); \
91#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
93 llvm_unreachable(
"unknown ACC data descriptor");
97 overrides[fn] = name.str();
101 if (
auto it = overrides.find(fn); it != overrides.end())
107 functionDisplayNameFn = std::move(fn);
112 if (functionDisplayNameFn)
113 return functionDisplayNameFn(mangledOrSymbol);
114 return mangledOrSymbol.str();
119 deviceTypeRuntimeValues[type] = runtimeValue;
123 if (
auto it = deviceTypeRuntimeValues.find(type);
124 it != deviceTypeRuntimeValues.end())
126 llvm::report_fatal_error(
127 llvm::Twine(
"missing OpenACC runtime device-type mapping for ") +
128 stringifyDeviceType(type));
133 mapFlagRuntimeValues[flag] = runtimeValue;
138 for (
unsigned bit = 0; bit != 32; ++bit) {
139 auto flag =
static_cast<MapFlags
>(1u << bit);
140 if (!bitEnumContainsAny(flags, flag))
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;
153 mapFlagsPostProcessFn = std::move(fn);
157 MapFlags flags)
const {
158 if (mapFlagsPostProcessFn)
159 return mapFlagsPostProcessFn(mapOp, flags);
165 std::string rendered = stringifyMapFlags(flags);
167 rendered += std::to_string(runtimeValue);
169 rendered += llvm::Twine::utohexstr(runtimeValue).str();
175 asyncSyncRuntimeValue = runtimeValue;
179 return asyncSyncRuntimeValue;
183 asyncNoValueRuntimeValue = runtimeValue;
187 return asyncNoValueRuntimeValue;
192 declareBinaryDescriptorFn = std::move(fn);
197 if (declareBinaryDescriptorFn)
198 return declareBinaryDescriptorFn(loc, builder);
199 return LLVM::ZeroOp::create(builder, loc,
200 LLVM::LLVMPointerType::get(builder.
getContext()));
205 for (uint32_t value = 0; value <= getMaxEnumValForDeviceType(); ++value)
206 if (std::optional<DeviceType> type = symbolizeDeviceType(value))
211 for (
unsigned bit = 0; bit != 32; ++bit) {
212 uint32_t value = 1u << bit;
213 if (std::optional<MapFlags> flag = symbolizeMapFlags(value))
218FailureOr<LLVM::CallOp>
225 StringRef symbolName = config.
getName(fn);
227 auto func = symbolTable.
lookup<LLVM::LLVMFuncOp>(symbolName);
231 if (
func.getFunctionType() != fnTy)
232 return emitError(loc) <<
"OpenACC runtime function '" << symbolName
233 <<
"' is already declared with signature "
234 <<
func.getFunctionType() <<
", expected " << fnTy;
238 func = LLVM::LLVMFuncOp::create(moduleBuilder, loc, symbolName, fnTy);
242 return LLVM::CallOp::create(builder, loc,
func, arguments);
MLIRContext * getContext() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
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...
Operation is the basic unit of execution within MLIR.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Configuration for OpenACC to LLVM runtime lowering.
int64_t getAsyncNoValueRuntimeValue() const
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
int64_t getAsyncSyncRuntimeValue() 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.