MLIR 24.0.0git
ACCToLLVMUtils.h
Go to the documentation of this file.
1//===- ACCToLLVMUtils.h - OpenACC to LLVM helpers ---------------*- 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#ifndef MLIR_CONVERSION_OPENACCTOLLVM_ACCTOLLVMUTILS_H
10#define MLIR_CONVERSION_OPENACCTOLLVM_ACCTOLLVMUTILS_H
11
14#include "mlir/IR/Builders.h"
15#include "mlir/IR/BuiltinOps.h"
16#include "mlir/IR/Location.h"
18#include "mlir/IR/Region.h"
20#include "llvm/ADT/STLExtras.h"
21#include "llvm/ADT/STLForwardCompat.h"
22#include "llvm/ADT/STLFunctionalExtras.h"
23#include "llvm/ADT/StringRef.h"
24
25#include <optional>
26#include <string>
27#include <utility>
28
29namespace mlir {
30namespace acc {
31
32/// Unfuses fused locations, returning the last sub-location.
33Location unfuseLoc(Location loc);
34
35/// Returns file:line:column location information when available.
36std::optional<FileLineColLoc> getFileLineColLoc(Location loc,
37 bool errorOnInvalidLocation);
38
39/// Returns the symbol name of the function \p op belongs to, or of \p op itself
40/// when it is a function.
41StringRef getParentFunctionName(Operation *op);
42
43/// Returns the enclosing function symbol name for \p value's defining op.
44StringRef getParentFunctionName(Value value);
45
46/// Returns the first non-empty enclosing function name from \p values.
47StringRef getParentFunctionName(ValueRange values);
48
49/// Returns the name to give a global that the conversion creates to hold
50/// \p detail of \p kind, such as the name of a variable or a source position.
51/// The dots make a name no identifier of the program can carry, so these
52/// globals are reachable by name without walking the symbols of the module.
53/// \p detail keeps letters, digits, `_`, `$` and `.` and folds every other
54/// character to an underscore. Distinct details can therefore collide on the
55/// same name; getOrCreateGlobalString is what then keeps the globals apart.
56std::string getInternalGlobalName(StringRef kind, StringRef detail);
57
58/// Creates or reuses a null-terminated string global in \p globalSymbolRegion.
59/// With \p symbolTable, a global already going by \p name is reused when it
60/// holds \p value, and a name that two different values arrive under gets a
61/// suffix to tell the globals apart. Without it the global is created under
62/// \p name as given, which only a caller whose names are unique by
63/// construction can ask for.
64Value getOrCreateGlobalString(Location loc, OpBuilder &builder, StringRef name,
65 StringRef value, Region &globalSymbolRegion,
66 SymbolTable *symbolTable = nullptr);
67
68/// Returns a pointer to a constant global holding an ident_t for OpenACC
69/// runtime calls. \p globalSymbolRegion and \p symbolTable are as in
70/// getOrCreateGlobalString; the ident and the source string it points to are
71/// named after the position they describe, so leaving out the table creates a
72/// set of them per call.
73Value createIdent(Location loc, StringRef functionName, OpBuilder &builder,
74 Region &globalSymbolRegion,
75 const ACCRuntimeCallConfig &config,
76 SymbolTable *symbolTable = nullptr);
77
78/// Sign-extends or truncates \p value to the i64 the runtime entry points take
79/// for values like queue numbers.
80Value castToI64(Location loc, Value value, OpBuilder &builder);
81
82/// Returns the queue an `async` clause selects: the value of the clause when it
83/// has one, the queue standing for an `async` clause without a value when
84/// \p asyncOnly is set, and the synchronous queue when there is no clause at
85/// all. \p asyncOperand must already be converted to the LLVM dialect.
86Value getAsyncQueue(Location loc, Value asyncOperand, bool asyncOnly,
87 OpBuilder &builder, const ACCRuntimeCallConfig &config);
88
89/// Emits the runtime call that waits for \p waitOperands on \p asyncQueue,
90/// which is what a `wait` clause or an `acc.wait` directive asks for. An empty
91/// \p waitOperands waits for every queue, as a `wait` clause without values
92/// does. The values must already be converted to the LLVM dialect.
93/// If provided, \p deviceNum must be an i32 value; otherwise the device number
94/// defaults to zero.
95LogicalResult emitWaitCall(Location loc, ValueRange waitOperands,
96 Value asyncQueue, OpBuilder &builder,
97 Region &globalSymbolRegion, SymbolTable &symbolTable,
98 const ACCRuntimeCallConfig &config,
99 Value deviceNum = {});
100
101/// Emits the runtime call that asks for the device address the object of the
102/// data clause \p clauseOp is mapped to, which is what a `use_device` clause
103/// exposes in the body of its construct. \p hostPtr is the address of that
104/// object, already converted to the LLVM dialect. With \p ifPresent, an object
105/// that is not mapped keeps its host address instead of being reported.
106///
107/// Bounds on the clause say which part of the object it names, which the
108/// address asked about does not have to state: the result stands for the
109/// object, so whatever reads it addresses the part it wants as it would on the
110/// host.
111FailureOr<Value> emitGetDevicePtrCall(Operation *clauseOp, Value hostPtr,
112 bool ifPresent, OpBuilder &builder,
113 Region &globalSymbolRegion,
114 SymbolTable &symbolTable,
115 const ACCRuntimeCallConfig &config);
116
117/// Runs \p emitFn guarded by a branch on \p ifCond, or unguarded when there is
118/// no condition. Leaves the insertion point after the guarded code, so that a
119/// caller can keep emitting into the same block either way.
120LogicalResult emitGuardedByIfCond(Location loc, Value ifCond,
121 RewriterBase &rewriter,
122 function_ref<LogicalResult()> emitFn);
123
124/// Emits \p thenFn on the path a branch on \p ifCond takes and \p elseFn on
125/// the other one, and returns the value that reaches the code following the
126/// branch. Both have to produce a value, and of the same type. With no
127/// condition only \p thenFn is emitted and its value returned. As with
128/// emitGuardedByIfCond, the insertion point is left after the branch.
129FailureOr<Value>
130emitValueSelectedByIfCond(Location loc, Value ifCond, RewriterBase &rewriter,
131 function_ref<FailureOr<Value>()> thenFn,
132 function_ref<FailureOr<Value>()> elseFn);
133
134/// Splices \p region, the body of a structured construct, into the block
135/// holding \p op, so that the construct itself can be erased. The code that
136/// follows \p op is branched to where the region ends, whether it ends in an
137/// `acc.terminator` or, as a region holding no terminator does, at the end of
138/// its blocks.
139void spliceConstructRegion(Operation *op, Region &region,
140 RewriterBase &rewriter);
141
142/// The clauses of a construct can be given once per device type. Of the values
143/// that reach a given device type, the ones naming it are the most specific,
144/// then the ones naming every device type, then the ones given before any
145/// device_type clause.
146SmallVector<DeviceType, 3> getDeviceTypesByPrecedence(DeviceType deviceType);
147
148namespace detail {
149/// The constructs carrying a `device_type` clause hold their async and wait
150/// clauses per device type. The others, such as `acc.enter_data` or
151/// `acc.kernel_environment`, hold a single value for each clause, spelling the
152/// value-less form either `asyncOnly`/`waitOnly` or `async`/`wait`.
153template <typename OpTy>
155 decltype(std::declval<OpTy>().hasAsyncOnly(DeviceType::None));
156template <typename OpTy>
157using has_async_only_t = decltype(std::declval<OpTy>().getAsyncOnly());
158template <typename OpTy>
159using has_wait_only_t = decltype(std::declval<OpTy>().getWaitOnly());
160} // namespace detail
161
162/// Returns the value the `async` clause of \p op names for \p deviceType, and
163/// sets \p asyncOnly when the clause names no queue. The value is the one the
164/// operation holds, so a caller in a conversion has to remap it.
165template <typename OpTy>
166Value getAsyncClauseValue(OpTy op, DeviceType deviceType, bool &asyncOnly) {
167 asyncOnly = false;
168 if constexpr (llvm::is_detected<detail::has_device_type_clauses_t,
169 OpTy>::value) {
170 for (DeviceType candidate : getDeviceTypesByPrecedence(deviceType)) {
171 if (op.hasAsyncOnly(candidate)) {
172 asyncOnly = true;
173 return {};
174 }
175 if (Value asyncValue = op.getAsyncValue(candidate))
176 return asyncValue;
177 }
178 return {};
179 } else if constexpr (llvm::is_detected<detail::has_async_only_t,
180 OpTy>::value) {
181 asyncOnly = op.getAsyncOnly();
182 return op.getAsyncOperand();
183 } else {
184 asyncOnly = op.getAsync();
185 return op.getAsyncOperand();
186 }
187}
188
189/// Appends to \p waitValues the queues the `wait` clause of \p op names for
190/// \p deviceType, and returns the device type the clause naming them is given
191/// for - a clause naming no queue waits for every one of them. Returns
192/// std::nullopt when \p op gives no such clause for \p deviceType. The values
193/// are the ones the operation holds, so a caller in a conversion has to remap
194/// them.
195template <typename OpTy>
196std::optional<DeviceType>
197getWaitClauseValues(OpTy op, DeviceType deviceType,
198 SmallVectorImpl<Value> &waitValues) {
199 if constexpr (llvm::is_detected<detail::has_device_type_clauses_t,
200 OpTy>::value) {
201 for (DeviceType candidate : getDeviceTypesByPrecedence(deviceType)) {
202 if (op.hasWaitOnly(candidate))
203 return candidate;
204 auto values = op.getWaitValues(candidate);
205 if (!values.empty()) {
206 llvm::append_range(waitValues, values);
207 return candidate;
208 }
209 }
210 return std::nullopt;
211 } else {
212 llvm::append_range(waitValues, op.getWaitOperands());
213 bool waitsForEveryQueue;
214 if constexpr (llvm::is_detected<detail::has_wait_only_t, OpTy>::value)
215 waitsForEveryQueue = op.getWaitOnly();
216 else
217 waitsForEveryQueue = op.getWait();
218 if (!waitsForEveryQueue && waitValues.empty())
219 return std::nullopt;
220 // The operation holds a single wait clause, which no device_type clause
221 // narrows to a device type.
222 return DeviceType::None;
223 }
224}
225
226/// Returns whether the `wait` clause \p op gives for \p deviceType carries a
227/// devnum modifier, which selects the device the queues belong to. Only the
228/// clause given for that device type is asked, so a caller passes the device
229/// type getWaitClauseValues took the queues from.
230template <typename OpTy>
231bool hasWaitDevnum(OpTy op, DeviceType deviceType) {
232 if constexpr (llvm::is_detected<detail::has_device_type_clauses_t,
233 OpTy>::value) {
234 return static_cast<bool>(op.getWaitDevnum(deviceType));
235 } else {
236 return static_cast<bool>(op.getWaitDevnum());
237 }
238}
239
240/// Returns the queue that the runtime calls of \p op run on, from the `async`
241/// clause it gives for \p deviceType.
242template <typename OpTy>
243Value getAsyncQueue(OpTy op, DeviceType deviceType,
244 ConversionPatternRewriter &rewriter,
245 const ACCRuntimeCallConfig &config) {
246 bool asyncOnly = false;
247 Value asyncValue = getAsyncClauseValue(op, deviceType, asyncOnly);
248 if (asyncValue)
249 asyncValue = rewriter.getRemappedValue(asyncValue);
250 return getAsyncQueue(op.getLoc(), asyncValue, asyncOnly, rewriter, config);
251}
252
253/// Emits the wait that a `wait` clause on \p op asks for before the runtime
254/// calls of the construct, waiting on \p asyncQueue for the queues the clause
255/// names, or for every queue when it names none. Nothing is emitted when there
256/// is no such clause for \p deviceType. A devnum modifier, which selects the
257/// device the queues belong to, is reported through \p accSupport as not yet
258/// implemented.
259template <typename OpTy>
260LogicalResult
261emitWaitClause(OpTy op, DeviceType deviceType, Value asyncQueue,
262 ConversionPatternRewriter &rewriter, OpenACCSupport &accSupport,
263 Region &globalSymbolRegion, SymbolTable &symbolTable,
264 const ACCRuntimeCallConfig &config) {
265 SmallVector<Value> waitValues;
266 std::optional<DeviceType> clauseDeviceType =
267 getWaitClauseValues(op, deviceType, waitValues);
268 if (!clauseDeviceType)
269 return success();
270 if (hasWaitDevnum(op, *clauseDeviceType)) {
271 (void)accSupport.emitNYI(op.getLoc(), "wait clause with a devnum modifier");
272 return failure();
273 }
274 for (Value &waitValue : waitValues)
275 waitValue = rewriter.getRemappedValue(waitValue);
276 return emitWaitCall(op.getLoc(), waitValues, asyncQueue, rewriter,
277 globalSymbolRegion, symbolTable, config);
278}
279
280} // namespace acc
281} // namespace mlir
282
283#endif // MLIR_CONVERSION_OPENACCTOLLVM_ACCTOLLVMUTILS_H
return success()
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
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.
InFlightDiagnostic emitNYI(Location loc, const Twine &message)
Report a case that is not yet supported by the implementation.
decltype(std::declval< OpTy >().getAsyncOnly()) has_async_only_t
decltype(std::declval< OpTy >().hasAsyncOnly(DeviceType::None)) has_device_type_clauses_t
The constructs carrying a device_type clause hold their async and wait clauses per device type.
decltype(std::declval< OpTy >().getWaitOnly()) has_wait_only_t
SmallVector< DeviceType, 3 > getDeviceTypesByPrecedence(DeviceType deviceType)
The clauses of a construct can be given once per device type.
std::optional< FileLineColLoc > getFileLineColLoc(Location loc, bool errorOnInvalidLocation)
Returns file:line:column location information when available.
std::optional< DeviceType > getWaitClauseValues(OpTy op, DeviceType deviceType, SmallVectorImpl< Value > &waitValues)
Appends to waitValues the queues the wait clause of op names for deviceType, and returns the device t...
LogicalResult emitWaitCall(Location loc, ValueRange waitOperands, Value asyncQueue, OpBuilder &builder, Region &globalSymbolRegion, SymbolTable &symbolTable, const ACCRuntimeCallConfig &config, Value deviceNum={})
Emits the runtime call that waits for waitOperands on asyncQueue, which is what a wait clause or an a...
Value getAsyncClauseValue(OpTy op, DeviceType deviceType, bool &asyncOnly)
Returns the value the async clause of op names for deviceType, and sets asyncOnly when the clause nam...
LogicalResult emitWaitClause(OpTy op, DeviceType deviceType, Value asyncQueue, ConversionPatternRewriter &rewriter, OpenACCSupport &accSupport, Region &globalSymbolRegion, SymbolTable &symbolTable, const ACCRuntimeCallConfig &config)
Emits the wait that a wait clause on op asks for before the runtime calls of the construct,...
Value getAsyncQueue(Location loc, Value asyncOperand, bool asyncOnly, OpBuilder &builder, const ACCRuntimeCallConfig &config)
Returns the queue an async clause selects: the value of the clause when it has one,...
bool hasWaitDevnum(OpTy op, DeviceType deviceType)
Returns whether the wait clause op gives for deviceType carries a devnum modifier,...
Value createIdent(Location loc, StringRef functionName, OpBuilder &builder, Region &globalSymbolRegion, const ACCRuntimeCallConfig &config, SymbolTable *symbolTable=nullptr)
Returns a pointer to a constant global holding an ident_t for OpenACC runtime calls.
FailureOr< Value > emitValueSelectedByIfCond(Location loc, Value ifCond, RewriterBase &rewriter, function_ref< FailureOr< Value >()> thenFn, function_ref< FailureOr< Value >()> elseFn)
Emits thenFn on the path a branch on ifCond takes and elseFn on the other one, and returns the value ...
std::string getInternalGlobalName(StringRef kind, StringRef detail)
Returns the name to give a global that the conversion creates to hold detail of kind,...
StringRef getParentFunctionName(Operation *op)
Returns the symbol name of the function op belongs to, or of op itself when it is a function.
LogicalResult emitGuardedByIfCond(Location loc, Value ifCond, RewriterBase &rewriter, function_ref< LogicalResult()> emitFn)
Runs emitFn guarded by a branch on ifCond, or unguarded when there is no condition.
Value castToI64(Location loc, Value value, OpBuilder &builder)
Sign-extends or truncates value to the i64 the runtime entry points take for values like queue number...
Value getOrCreateGlobalString(Location loc, OpBuilder &builder, StringRef name, StringRef value, Region &globalSymbolRegion, SymbolTable *symbolTable=nullptr)
Creates or reuses a null-terminated string global in globalSymbolRegion.
void spliceConstructRegion(Operation *op, Region &region, RewriterBase &rewriter)
Splices region, the body of a structured construct, into the block holding op, so that the construct ...
Location unfuseLoc(Location loc)
Unfuses fused locations, returning the last sub-location.
FailureOr< Value > emitGetDevicePtrCall(Operation *clauseOp, Value hostPtr, bool ifPresent, OpBuilder &builder, Region &globalSymbolRegion, SymbolTable &symbolTable, const ACCRuntimeCallConfig &config)
Emits the runtime call that asks for the device address the object of the data clause clauseOp is map...
Include the generated interface declarations.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147