22#include "llvm/ADT/STLExtras.h"
31 ConversionPatternRewriter &rewriter) {
32 Type i32Ty = rewriter.getI32Type();
35 return arith::TruncIOp::create(rewriter, loc, i32Ty, value);
38 return arith::ExtUIOp::create(rewriter, loc, i32Ty, value);
39 return arith::ExtSIOp::create(rewriter, loc, i32Ty, value);
44template <
typename 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),
55 Region &globalSymbolRegion;
56 SymbolTable &symbolTable;
57 ACCRuntimeCallConfig config;
60struct WaitOpLowering :
public ACCExecutableDirectivePattern<WaitOp> {
61 using ACCExecutableDirectivePattern<WaitOp>::ACCExecutableDirectivePattern;
64 matchAndRewrite(WaitOp op, WaitOp::Adaptor adaptor,
65 ConversionPatternRewriter &rewriter)
const override {
66 Location loc = op->getLoc();
68 auto emitWait = [&]() -> LogicalResult {
69 Value asyncOperand = op.getAsyncOperand()
70 ? rewriter.getRemappedValue(op.getAsyncOperand())
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));
78 Value deviceNum = adaptor.getWaitDevnum();
81 bool isUnsigned = op.getWaitDevnum().getType().isUnsignedInteger();
82 deviceNum = castToI32(loc, deviceNum, isUnsigned, rewriter);
85 return emitWaitCall(loc, waitValues, asyncQueue, rewriter,
86 globalSymbolRegion, symbolTable, config, deviceNum);
102 Value deviceNum, StringRef functionName,
104 ConversionPatternRewriter &rewriter,
106 Type i64Ty = rewriter.getI64Type();
107 Value deviceTypeValue = LLVM::ConstantOp::create(
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);
117 {ident, flags, deviceTypeValue, deviceNumValue});
120static LogicalResult rewriteInitOrShutdown(
Operation *op,
Value deviceNum,
122 Value ifCond,
bool isInit,
123 ConversionPatternRewriter &rewriter,
124 Region &globalSymbolRegion,
129 auto emitCalls = [&]() -> LogicalResult {
133 : RuntimeFunction::ACCRTL_tgt_acc_shutdown;
135 auto emitOne = [&](DeviceType deviceType) {
136 return emitDeviceOperationCall(loc, fn, deviceType, deviceNum,
137 functionName, globalSymbolRegion,
138 symbolTable, rewriter, config);
141 if (!deviceTypesAttr)
142 return emitOne(DeviceType::None);
145 if (
auto typeAttr = dyn_cast<DeviceTypeAttr>(attr))
146 if (
failed(emitOne(typeAttr.getValue())))
155 rewriter.eraseOp(op);
159struct InitOpLowering :
public ACCExecutableDirectivePattern<InitOp> {
160 using ACCExecutableDirectivePattern<InitOp>::ACCExecutableDirectivePattern;
163 matchAndRewrite(InitOp op, InitOp::Adaptor adaptor,
164 ConversionPatternRewriter &rewriter)
const override {
165 return rewriteInitOrShutdown(
166 op, adaptor.getDeviceNum(), op.getDeviceTypesAttr(), op.getIfCond(),
167 true, rewriter, globalSymbolRegion, symbolTable, config);
171struct ShutdownOpLowering :
public ACCExecutableDirectivePattern<ShutdownOp> {
172 using ACCExecutableDirectivePattern<
173 ShutdownOp>::ACCExecutableDirectivePattern;
176 matchAndRewrite(ShutdownOp op, ShutdownOp::Adaptor adaptor,
177 ConversionPatternRewriter &rewriter)
const override {
178 return rewriteInitOrShutdown(
179 op, adaptor.getDeviceNum(), op.getDeviceTypesAttr(), op.getIfCond(),
180 false, rewriter, globalSymbolRegion, symbolTable, config);
184struct SetOpLowering :
public ACCExecutableDirectivePattern<SetOp> {
185 using ACCExecutableDirectivePattern<SetOp>::ACCExecutableDirectivePattern;
188 matchAndRewrite(SetOp op, SetOp::Adaptor adaptor,
189 ConversionPatternRewriter &rewriter)
const override {
190 Location loc = op.getLoc();
191 Type i64Ty = rewriter.getI64Type();
193 auto emitSet = [&]() -> LogicalResult {
194 if (Value asyncValue = adaptor.getDefaultAsync()) {
195 asyncValue =
castToI64(loc, asyncValue, rewriter);
198 globalSymbolRegion, config, &symbolTable);
200 loc, rewriter, globalSymbolRegion, symbolTable,
201 RuntimeFunction::ACCRTL_tgt_acc_set_default_async, config,
202 {ident, asyncValue})))
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,
214 symbolTable, rewriter, config);
216 if (
auto deviceTypeAttr = op.getDeviceTypeAttr()) {
217 Value deviceTypeValue = LLVM::ConstantOp::create(
218 rewriter, loc, i64Ty,
220 Value ident =
createIdent(loc, StringRef(), rewriter,
221 globalSymbolRegion, config, &symbolTable);
222 Value flags = LLVM::ConstantOp::create(rewriter, loc, i64Ty, 0);
224 loc, rewriter, globalSymbolRegion, symbolTable,
225 RuntimeFunction::ACCRTL_tgt_acc_set_device_type, config,
226 {ident, flags, deviceTypeValue});
234 rewriter.eraseOp(op);
243 target.addIllegalOp<acc::InitOp, acc::ShutdownOp, acc::WaitOp, acc::SetOp>();
251 .
add<WaitOpLowering, InitOpLowering, ShutdownOpLowering, SetOpLowering>(
252 converter, globalSymbolRegion, symbolTable, config);
Attributes are known-constant values of operations.
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
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...
Operation is the basic unit of execution within MLIR.
Location getLoc()
The source location the operation was defined or derived from.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
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...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
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 ...
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.