MLIR 24.0.0git
OpenACCRuntimeUtils.h
Go to the documentation of this file.
1//===- OpenACCRuntimeUtils.h - 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//
9// Utilities for resolving OpenACC compiler-to-runtime entry points declared in
10// OpenACCRuntimeFunctions.def.
11//
12//===----------------------------------------------------------------------===//
13
14#ifndef MLIR_DIALECT_OPENACC_OPENACCRUNTIMEUTILS_H
15#define MLIR_DIALECT_OPENACC_OPENACCRUNTIMEUTILS_H
16
19#include "mlir/IR/Builders.h"
20#include "mlir/IR/BuiltinOps.h"
21#include "mlir/IR/Region.h"
23#include "llvm/ADT/ArrayRef.h"
24#include "llvm/ADT/DenseMap.h"
25#include "llvm/ADT/StringRef.h"
26
27#include <cstdint>
28#include <functional>
29#include <string>
30
31namespace mlir {
32namespace acc {
33
34/// IDs for OpenACC compiler-to-runtime entry points (`__tgt_acc_*`).
35enum class RuntimeFunction {
36#define ACC_RTL(Enum, ...) Enum,
37#include "mlir/Dialect/OpenACC/OpenACCRuntimeFunctions.def"
38};
39
40/// Returns the default runtime symbol name for \p fn.
42
43/// Builds the LLVM function type for \p fn in \p ctx.
44LLVM::LLVMFunctionType getRuntimeFunctionType(MLIRContext *ctx,
46
47/// IDs for the argument descriptors of the OpenACC data entry points, declared
48/// in OpenACCRuntimeDescriptors.def.
49enum class DataDescriptor {
50#define ACC_DESC_BEGIN(Enum, ...) Enum,
51#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
52};
53
54/// Field indices of each descriptor, in declaration order.
55#define ACC_DESC_BEGIN(Enum, ...) enum class Enum##Field : int64_t {
56#define ACC_DESC_FIELD(Enum, Field, TypeToken) Field,
57#define ACC_DESC_END(Enum) \
58 } \
59 ;
60#include "mlir/Dialect/OpenACC/OpenACCRuntimeDescriptors.def"
61
62/// The index an insert or extract of \p field addresses.
63template <typename FieldEnum>
65 return static_cast<int64_t>(field);
66}
67
68/// Returns the name of the runtime type \p desc materializes.
70
71/// Returns the descriptor kind the runtime reads from the version field of
72/// \p desc. An overlay contributes its kind to the descriptor it nests.
73DataDescKind getDataDescriptorKind(DataDescriptor desc);
74
75/// Builds the LLVM type of \p desc in \p ctx. A descriptor that nests another
76/// one takes it as \p baseType, which is ignored otherwise.
77LLVM::LLVMStructType getDataDescriptorType(MLIRContext *ctx,
78 DataDescriptor desc,
79 Type baseType = {});
80
81/// Configuration for OpenACC to LLVM runtime lowering. Device types and map
82/// flags use the OpenACC dialect encodings by default.
84public:
85 using FunctionDisplayNameFn = std::function<std::string(StringRef)>;
86
88
89 void setName(RuntimeFunction fn, StringRef name);
90 StringRef getName(RuntimeFunction fn) const;
91
93 std::string getFunctionDisplayName(StringRef mangledOrSymbol) const;
94
95 /// Map an OpenACC dialect \p DeviceType to the integer encoding expected by
96 /// the target runtime. Dialect ordinals and runtime ABI values are not
97 /// required to match. A target whose ABI differs from the default dialect
98 /// encoding must install its mapping here. Querying an unmapped type is an
99 /// error.
100 void setDeviceTypeRuntimeValue(DeviceType type, int64_t runtimeValue);
101 int64_t getDeviceTypeRuntimeValue(DeviceType type) const;
102
103 /// Map a single OpenACC dialect \p MapFlags bit to the bit the target runtime
104 /// gives the same meaning. getMapFlagsRuntimeValue combines the bits of a set
105 /// of flags into the encoding the runtime reads from an argument-type slot.
106 /// As with device types, dialect and runtime encodings are not required to
107 /// match. A target whose ABI differs from the default dialect encoding must
108 /// install its mapping here, and a set bit with no mapping is an error.
109 void setMapFlagRuntimeValue(MapFlags flag, int64_t runtimeValue);
110 int64_t getMapFlagsRuntimeValue(MapFlags flags) const;
111
112 /// Adjust the flags of a mapping before they are encoded. The hook is keyed
113 /// on the operation that states the mapping and is applied to every mapped
114 /// object before anything is derived from the flags. The flags are used as
115 /// computed when no hook is installed.
117 std::function<MapFlags(Operation *mapOp, MapFlags flags)>;
119 MapFlags postProcessMapFlags(Operation *mapOp, MapFlags flags) const;
120
121 /// Renders \p flags for a diagnostic as the names of the set bits and the
122 /// decimal and hexadecimal encoding the runtime reads.
123 std::string formatMapFlags(MapFlags flags) const;
124
125 /// Runtime encoding of `acc_async_sync`, used when an operation carries no
126 /// `async` clause. OpenACC defines the name of this queue but leaves its
127 /// value to the implementation, so it is part of the runtime ABI.
128 void setAsyncSyncRuntimeValue(int64_t runtimeValue);
130
131 /// Runtime encoding of `acc_async_noval`, used for an `async` clause without
132 /// an argument. As with `acc_async_sync`, the value is implementation-defined
133 void setAsyncNoValueRuntimeValue(int64_t runtimeValue);
135
136 /// Materialize the target-specific binary descriptor passed to
137 /// `__tgt_acc_declare`. The default is a null pointer.
141
142private:
144 DenseMap<DeviceType, int64_t> deviceTypeRuntimeValues;
145 DenseMap<MapFlags, int64_t> mapFlagRuntimeValues;
146 FunctionDisplayNameFn functionDisplayNameFn;
147 MapFlagsPostProcessFn mapFlagsPostProcessFn;
148 DeclareBinaryDescriptorFn declareBinaryDescriptorFn;
149 // Default to the encodings used by openacc.h (`acc_async_sync` /
150 // `acc_async_noval`).
151 int64_t asyncSyncRuntimeValue = -1;
152 int64_t asyncNoValueRuntimeValue = -4;
153};
154
155/// Install a device-type mapping that uses OpenACC dialect enum ordinals as the
156/// runtime encoding. This is only correct when the target runtime happens to
157/// use the same numbering; runtimes with a different ABI must install their
158/// own mapping via \c setDeviceTypeRuntimeValue.
160
161/// Install a map-flag mapping that uses the OpenACC dialect bit positions as
162/// the runtime encoding, with the same caveat as
163/// \c populateDialectIdentityDeviceTypeMapping.
165
166/// Declares (if needed) and returns a call to the runtime function identified
167/// by \p fn using the name from \p config. Fails and emits a diagnostic if the
168/// symbol is already declared with a signature the runtime cannot be called
169/// through. The declaration is created in \p globalSymbolRegion and
170/// registered in \p symbolTable.
171FailureOr<LLVM::CallOp> createRuntimeCall(Location loc, OpBuilder &builder,
172 Region &globalSymbolRegion,
173 SymbolTable &symbolTable,
175 const ACCRuntimeCallConfig &config,
176 ArrayRef<Value> arguments);
177
178} // namespace acc
179} // namespace mlir
180
181#endif // MLIR_DIALECT_OPENACC_OPENACCRUNTIMEUTILS_H
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
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
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
Definition SymbolTable.h:25
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,...
int64_t getDataDescriptorFieldIndex(FieldEnum field)
The index an insert or extract of field addresses.
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.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120