MLIR 24.0.0git
ACCHostDataPatterns.cpp
Go to the documentation of this file.
1//===- ACCHostDataPatterns.cpp - ACC host_data to LLVM ---------*- 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// Lowering of the OpenACC host_data construct to the OpenACC runtime. Each
10// use_device clause of the construct becomes the call that asks the runtime
11// for the device address its object is mapped to, so that the body of the
12// construct works on device addresses. An if clause on the construct decides
13// between that address and the host one, and the body is left in place of the
14// construct.
15//
16//===----------------------------------------------------------------------===//
17
20
24
25using namespace mlir;
26using namespace mlir::acc;
27
28/// Returns the host_data construct \p op gives a device address to, or null
29/// when the clause belongs to no such construct.
30static HostDataOp getHostDataConstruct(UseDeviceOp op) {
31 for (Operation *user : op->getUsers())
32 if (auto hostData = dyn_cast<HostDataOp>(user))
33 return hostData;
34 return nullptr;
35}
36
37namespace {
38/// A use_device clause stands for the device address of the object it names,
39/// so it becomes the call that asks the runtime for that address. The call is
40/// emitted before the construct, so that every use of the clause in its body
41/// is reached by it.
42struct UseDeviceOpLowering : public ConvertOpToLLVMPattern<UseDeviceOp> {
43 UseDeviceOpLowering(const LLVMTypeConverter &converter,
44 Region &globalSymbolRegion, SymbolTable &symbolTable,
45 const ACCRuntimeCallConfig &config,
46 PatternBenefit benefit = 1)
47 : ConvertOpToLLVMPattern<UseDeviceOp>(converter, benefit),
48 globalSymbolRegion(globalSymbolRegion), symbolTable(symbolTable),
49 config(config) {}
50
51 LogicalResult
52 matchAndRewrite(UseDeviceOp op, OpAdaptor adaptor,
53 ConversionPatternRewriter &rewriter) const override {
54 HostDataOp hostData = getHostDataConstruct(op);
55 if (!hostData)
56 return rewriter.notifyMatchFailure(
57 op, "use_device clause of no host_data construct");
58
59 Location loc = op.getLoc();
60 Value hostPtr = adaptor.getVar();
61 auto emitCall = [&]() -> FailureOr<Value> {
62 return acc::emitGetDevicePtrCall(op, hostPtr, hostData.getIfPresent(),
63 rewriter, globalSymbolRegion,
64 symbolTable, config);
65 };
66 // An if clause that does not hold leaves the body on host addresses.
67 auto keepHostPtr = [&]() -> FailureOr<Value> { return hostPtr; };
68
69 Value ifCond = hostData.getIfCond()
70 ? rewriter.getRemappedValue(hostData.getIfCond())
71 : Value();
72 if (ifCond)
73 rewriter.setInsertionPoint(hostData);
74 FailureOr<Value> devicePtr = acc::emitValueSelectedByIfCond(
75 loc, ifCond, rewriter, emitCall, keepHostPtr);
76 if (failed(devicePtr))
77 return failure();
78
79 rewriter.replaceOp(op, *devicePtr);
80 return success();
81 }
82
83 Region &globalSymbolRegion;
84 SymbolTable &symbolTable;
85 ACCRuntimeCallConfig config;
86};
87
88/// The construct itself has nothing left to do once its clauses hold device
89/// addresses, so its body takes its place.
90struct HostDataOpLowering : public ConvertOpToLLVMPattern<HostDataOp> {
91 using ConvertOpToLLVMPattern<HostDataOp>::ConvertOpToLLVMPattern;
92
93 LogicalResult
94 matchAndRewrite(HostDataOp op, OpAdaptor,
95 ConversionPatternRewriter &rewriter) const override {
96 acc::spliceConstructRegion(op, op.getRegion(), rewriter);
97 rewriter.eraseOp(op);
98 return success();
99 }
100};
101
102} // namespace
103
105 target.addIllegalOp<acc::HostDataOp, acc::UseDeviceOp>();
106}
107
109 LLVMTypeConverter &converter, RewritePatternSet &patterns,
110 Region &globalSymbolRegion, SymbolTable &symbolTable,
111 const acc::ACCRuntimeCallConfig &config) {
112 patterns.add<UseDeviceOpLowering>(converter, globalSymbolRegion, symbolTable,
113 config);
114 patterns.add<HostDataOpLowering>(converter);
115}
static HostDataOp getHostDataConstruct(UseDeviceOp op)
Returns the host_data construct op gives a device address to, or null when the clause belongs to no s...
return success()
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.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
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
Configuration for OpenACC to LLVM runtime lowering.
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 ...
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 ...
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...
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
void configureACCHostDataConversionLegality(ConversionTarget &target)
Configure conversion legality for the OpenACC host_data construct.
void populateACCHostDataPatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns, Region &globalSymbolRegion, SymbolTable &symbolTable, const acc::ACCRuntimeCallConfig &config={})
Populate patterns that lower the OpenACC host_data construct to runtime calls, one per use_device cla...