MLIR 24.0.0git
ACCExecutableDirectivePatterns.cpp
Go to the documentation of this file.
1//===- ACCExecutableDirectivePatterns.cpp - ACC exec patterns ---*- 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// Lowers OpenACC executable directives (init, shutdown, wait, set) to calls to
10// an OpenACC offloading runtime compiler interface.
11//
12//===----------------------------------------------------------------------===//
13
16
22#include "llvm/ADT/STLExtras.h"
23
24#include <cstdint>
25
26using namespace mlir;
27using namespace mlir::acc;
28
29namespace {
30static Value castToI32(Location loc, Value value, bool isUnsigned,
31 ConversionPatternRewriter &rewriter) {
32 Type i32Ty = rewriter.getI32Type();
33 unsigned bitwidth = value.getType().getIntOrFloatBitWidth();
34 if (bitwidth > 32)
35 return arith::TruncIOp::create(rewriter, loc, i32Ty, value);
36 if (bitwidth < 32) {
37 if (isUnsigned)
38 return arith::ExtUIOp::create(rewriter, loc, i32Ty, value);
39 return arith::ExtSIOp::create(rewriter, loc, i32Ty, value);
40 }
41 return value;
42}
43
44template <typename OpTy>
45struct ACCExecutableDirectivePattern : public ConvertOpToLLVMPattern<OpTy> {
46 ACCExecutableDirectivePattern(const LLVMTypeConverter &converter,
47 Region &globalSymbolRegion,
48 SymbolTable &symbolTable,
49 const ACCRuntimeCallConfig &config,
50 PatternBenefit benefit = 1)
51 : ConvertOpToLLVMPattern<OpTy>(converter, benefit),
52 globalSymbolRegion(globalSymbolRegion), symbolTable(symbolTable),
53 config(config) {}
54
55 Region &globalSymbolRegion;
56 SymbolTable &symbolTable;
57 ACCRuntimeCallConfig config;
58};
59
60struct WaitOpLowering : public ACCExecutableDirectivePattern<WaitOp> {
61 using ACCExecutableDirectivePattern<WaitOp>::ACCExecutableDirectivePattern;
62
63 LogicalResult
64 matchAndRewrite(WaitOp op, WaitOp::Adaptor adaptor,
65 ConversionPatternRewriter &rewriter) const override {
66 Location loc = op->getLoc();
67
68 auto emitWait = [&]() -> LogicalResult {
69 Value asyncOperand = op.getAsyncOperand()
70 ? rewriter.getRemappedValue(op.getAsyncOperand())
71 : Value();
72 Value asyncQueue =
73 getAsyncQueue(loc, asyncOperand, op.getAsync(), rewriter, config);
74 SmallVector<Value> waitValues;
75 for (Value operand : op.getWaitOperands())
76 waitValues.push_back(rewriter.getRemappedValue(operand));
77
78 Value deviceNum = adaptor.getWaitDevnum();
79 if (deviceNum) {
80 // The converted operand is signless, so inspect the original type.
81 bool isUnsigned = op.getWaitDevnum().getType().isUnsignedInteger();
82 deviceNum = castToI32(loc, deviceNum, isUnsigned, rewriter);
83 }
84
85 return emitWaitCall(loc, waitValues, asyncQueue, rewriter,
86 globalSymbolRegion, symbolTable, config, deviceNum);
87 };
88
89 if (failed(emitGuardedByIfCond(loc, op.getIfCond(), rewriter, emitWait)))
90 return failure();
91
92 rewriter.eraseOp(op);
93 return success();
94 }
95};
96
97/// Emit a call to a runtime entry point taking
98/// `(ident, flags, deviceType, deviceNum)`. A null `deviceNum` selects the
99/// current device.
100static LogicalResult
101emitDeviceOperationCall(Location loc, RuntimeFunction fn, DeviceType deviceType,
102 Value deviceNum, StringRef functionName,
103 Region &globalSymbolRegion, SymbolTable &symbolTable,
104 ConversionPatternRewriter &rewriter,
105 const ACCRuntimeCallConfig &config) {
106 Type i64Ty = rewriter.getI64Type();
107 Value deviceTypeValue = LLVM::ConstantOp::create(
108 rewriter, loc, i64Ty, config.getDeviceTypeRuntimeValue(deviceType));
109 Value ident = createIdent(loc, functionName, rewriter, globalSymbolRegion,
110 config, &symbolTable);
111 Value flags = LLVM::ConstantOp::create(rewriter, loc, i64Ty, 0);
112 Value deviceNumValue =
113 deviceNum ? castToI64(loc, deviceNum, rewriter)
114 : LLVM::ConstantOp::create(rewriter, loc, i64Ty, -1);
115 return createRuntimeCall(loc, rewriter, globalSymbolRegion, symbolTable, fn,
116 config,
117 {ident, flags, deviceTypeValue, deviceNumValue});
118}
119
120static LogicalResult rewriteInitOrShutdown(Operation *op, Value deviceNum,
121 ArrayAttr deviceTypesAttr,
122 Value ifCond, bool isInit,
123 ConversionPatternRewriter &rewriter,
124 Region &globalSymbolRegion,
125 SymbolTable &symbolTable,
126 const ACCRuntimeCallConfig &config) {
127 Location loc = op->getLoc();
128
129 auto emitCalls = [&]() -> LogicalResult {
130 StringRef functionName = deviceNum ? getParentFunctionName(deviceNum)
132 RuntimeFunction fn = isInit ? RuntimeFunction::ACCRTL_tgt_acc_init
133 : RuntimeFunction::ACCRTL_tgt_acc_shutdown;
134
135 auto emitOne = [&](DeviceType deviceType) {
136 return emitDeviceOperationCall(loc, fn, deviceType, deviceNum,
137 functionName, globalSymbolRegion,
138 symbolTable, rewriter, config);
139 };
140
141 if (!deviceTypesAttr)
142 return emitOne(DeviceType::None);
143
144 for (Attribute attr : deviceTypesAttr) {
145 if (auto typeAttr = dyn_cast<DeviceTypeAttr>(attr))
146 if (failed(emitOne(typeAttr.getValue())))
147 return failure();
148 }
149 return success();
150 };
151
152 if (failed(emitGuardedByIfCond(loc, ifCond, rewriter, emitCalls)))
153 return failure();
154
155 rewriter.eraseOp(op);
156 return success();
157}
158
159struct InitOpLowering : public ACCExecutableDirectivePattern<InitOp> {
160 using ACCExecutableDirectivePattern<InitOp>::ACCExecutableDirectivePattern;
161
162 LogicalResult
163 matchAndRewrite(InitOp op, InitOp::Adaptor adaptor,
164 ConversionPatternRewriter &rewriter) const override {
165 return rewriteInitOrShutdown(
166 op, adaptor.getDeviceNum(), op.getDeviceTypesAttr(), op.getIfCond(),
167 /*isInit=*/true, rewriter, globalSymbolRegion, symbolTable, config);
168 }
169};
170
171struct ShutdownOpLowering : public ACCExecutableDirectivePattern<ShutdownOp> {
172 using ACCExecutableDirectivePattern<
173 ShutdownOp>::ACCExecutableDirectivePattern;
174
175 LogicalResult
176 matchAndRewrite(ShutdownOp op, ShutdownOp::Adaptor adaptor,
177 ConversionPatternRewriter &rewriter) const override {
178 return rewriteInitOrShutdown(
179 op, adaptor.getDeviceNum(), op.getDeviceTypesAttr(), op.getIfCond(),
180 /*isInit=*/false, rewriter, globalSymbolRegion, symbolTable, config);
181 }
182};
183
184struct SetOpLowering : public ACCExecutableDirectivePattern<SetOp> {
185 using ACCExecutableDirectivePattern<SetOp>::ACCExecutableDirectivePattern;
186
187 LogicalResult
188 matchAndRewrite(SetOp op, SetOp::Adaptor adaptor,
189 ConversionPatternRewriter &rewriter) const override {
190 Location loc = op.getLoc();
191 Type i64Ty = rewriter.getI64Type();
192
193 auto emitSet = [&]() -> LogicalResult {
194 if (Value asyncValue = adaptor.getDefaultAsync()) {
195 asyncValue = castToI64(loc, asyncValue, rewriter);
196 Value ident =
197 createIdent(loc, getParentFunctionName(asyncValue), rewriter,
198 globalSymbolRegion, config, &symbolTable);
200 loc, rewriter, globalSymbolRegion, symbolTable,
201 RuntimeFunction::ACCRTL_tgt_acc_set_default_async, config,
202 {ident, asyncValue})))
203 return failure();
204 }
205
206 if (op.getDeviceNum()) {
207 Value deviceNum = adaptor.getDeviceNum();
208 DeviceType deviceType = DeviceType::None;
209 if (auto deviceTypeAttr = op.getDeviceTypeAttr())
210 deviceType = deviceTypeAttr.getValue();
211 return emitDeviceOperationCall(
212 loc, RuntimeFunction::ACCRTL_tgt_acc_set_device_num, deviceType,
213 deviceNum, getParentFunctionName(deviceNum), globalSymbolRegion,
214 symbolTable, rewriter, config);
215 }
216 if (auto deviceTypeAttr = op.getDeviceTypeAttr()) {
217 Value deviceTypeValue = LLVM::ConstantOp::create(
218 rewriter, loc, i64Ty,
219 config.getDeviceTypeRuntimeValue(deviceTypeAttr.getValue()));
220 Value ident = createIdent(loc, StringRef(), rewriter,
221 globalSymbolRegion, config, &symbolTable);
222 Value flags = LLVM::ConstantOp::create(rewriter, loc, i64Ty, 0);
223 return createRuntimeCall(
224 loc, rewriter, globalSymbolRegion, symbolTable,
225 RuntimeFunction::ACCRTL_tgt_acc_set_device_type, config,
226 {ident, flags, deviceTypeValue});
227 }
228 return success();
229 };
230
231 if (failed(emitGuardedByIfCond(loc, op.getIfCond(), rewriter, emitSet)))
232 return failure();
233
234 rewriter.eraseOp(op);
235 return success();
236 }
237};
238
239} // namespace
240
243 target.addIllegalOp<acc::InitOp, acc::ShutdownOp, acc::WaitOp, acc::SetOp>();
244}
245
247 LLVMTypeConverter &converter, RewritePatternSet &patterns,
248 Region &globalSymbolRegion, SymbolTable &symbolTable,
249 const acc::ACCRuntimeCallConfig &config) {
250 patterns
251 .add<WaitOpLowering, InitOpLowering, ShutdownOpLowering, SetOpLowering>(
252 converter, globalSymbolRegion, symbolTable, config);
253}
return success()
ArrayAttr()
Attributes are known-constant values of operations.
Definition Attributes.h:25
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
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
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Configuration for OpenACC to LLVM runtime lowering.
int64_t getDeviceTypeRuntimeValue(DeviceType type) const
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 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,...
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.
RuntimeFunction
IDs for OpenACC compiler-to-runtime entry points (__tgt_acc_*).
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...
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 ...
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
void populateACCExecutableDirectivePatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns, Region &globalSymbolRegion, SymbolTable &symbolTable, const acc::ACCRuntimeCallConfig &config={})
Populate patterns that lower OpenACC executable directives (init, shutdown, wait, set) to LLVM runtim...
void configureACCExecutableDirectiveConversionLegality(ConversionTarget &target)
Configure conversion legality for OpenACC executable directives lowered to runtime calls.