MLIR 24.0.0git
HostOpFiltering.cpp
Go to the documentation of this file.
1//===- HostOpFiltering.cpp ------------------------------------------------===//
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// This file implements transforms to swap stack allocations on the target
10// device with device shared memory where applicable.
11//
12//===----------------------------------------------------------------------===//
13
15
19
20namespace mlir {
21namespace omp {
22#define GEN_PASS_DEF_HOSTOPFILTERINGPASS
23#include "mlir/Dialect/OpenMP/Transforms/Passes.h.inc"
24} // namespace omp
25} // namespace mlir
26
27using namespace mlir;
28
29/// Some host operations, like \c llvm.mlir.addressof and constants, must remain
30/// in the device module because they impact how device code is generated when
31/// attached to an \c omp.target operation.
32///
33/// This function identifies the operations that need this special handling.
34/// This includes cast-style operations to avoid losing information about the
35/// original source of an operand.
36static bool keepHostOpInDevice(Operation &op) {
37 return isPure(&op) &&
38 op.getDialect() ==
39 op.getContext()->getLoadedDialect<LLVM::LLVMDialect>();
40}
41
42/// Add an \c omp.map.info operation and all its members recursively to the
43/// output set to be later rewritten.
44///
45/// Dependencies across \c omp.map.info are maintained by ensuring dependencies
46/// are added to the output sets before operations based on them.
47static void collectRewrite(omp::MapInfoOp mapOp,
49 for (Value member : mapOp.getMembers())
50 collectRewrite(cast<omp::MapInfoOp>(member.getDefiningOp()), rewrites);
51
52 rewrites.insert(mapOp);
53}
54
55/// Add the given value to a sorted set if it should be replaced by a
56/// placeholder when used as an operand that must remain for the device.
57///
58/// Values that are block arguments of function operations are skipped, since
59/// they will still be available after all rewrites are completed, and operands
60/// of operations that need to remain on the host are recursively collected.
61static void collectRewrite(Value value, llvm::SetVector<Value> &rewrites) {
62 if ((isa<BlockArgument>(value) &&
63 isa<FunctionOpInterface>(
64 cast<BlockArgument>(value).getOwner()->getParentOp())) ||
65 rewrites.contains(value))
66 return;
67
68 Operation *op = value.getDefiningOp();
69 if (op && keepHostOpInDevice(*op))
70 for (Value operand : op->getOperands())
71 collectRewrite(operand, rewrites);
72
73 rewrites.insert(value);
74}
75
76/// Provide the \c device_type of an \c omp.declare_target attribute, if
77/// defined.
78static std::optional<omp::DeclareTargetDeviceType>
80 auto declareTargetOp = dyn_cast<omp::DeclareTargetInterface>(op);
81 omp::DeclareTargetAttr declareTargetAttr =
82 declareTargetOp ? declareTargetOp.getDeclareTarget() : nullptr;
83 if (declareTargetAttr)
84 return declareTargetAttr.getDeviceType();
85 return std::nullopt;
86}
87
88namespace {
89class HostOpFilteringPass
90 : public omp::impl::HostOpFilteringPassBase<HostOpFilteringPass> {
91public:
92 HostOpFilteringPass() = default;
93
94 void runOnOperation() override {
95 auto op = dyn_cast<omp::OffloadModuleInterface>(getOperation());
96 if (!op || !op.getIsTargetDevice())
97 return;
98
99 op->walk<WalkOrder::PreOrder>([&](LLVM::LLVMFuncOp funcOp) {
100 omp::DeclareTargetDeviceType declareType =
101 getDeclareTargetDevice(*funcOp.getOperation())
102 .value_or(omp::DeclareTargetDeviceType::host);
103
104 // Only process host function definitions.
105 if (funcOp.isExternal() ||
106 declareType != omp::DeclareTargetDeviceType::host)
107 return WalkResult::advance();
108
109 if (failed(rewriteHostFunction(funcOp))) {
110 funcOp.emitOpError() << "could not filter host-only operations";
111 return WalkResult::interrupt();
112 }
113 return WalkResult::advance();
114 });
115
116 // Make non-declare target globals internal for the device. They cannot be
117 // deleted, because they are needed in order to properly lower map clauses.
118 // However, no uses will remain in the device module, so we make them
119 // internal to prevent link time redefinitions.
120 op->walk([&](LLVM::GlobalOp globalOp) {
121 if (!getDeclareTargetDevice(*globalOp.getOperation()).has_value())
122 globalOp.setLinkage(LLVM::Linkage::Internal);
123 });
124 }
125
126private:
127 /// Rewrite the given host device function containing \c omp.target
128 /// operations, to remove host-only operations that are not used by device
129 /// codegen.
130 ///
131 /// It is based on the expected form of an MLIR module lowered to where it can
132 /// be directly translated to LLVM IR and it performs the following mutations:
133 /// - Removes all returned values from the function.
134 /// - \c omp.target operations are moved to the end of the function. If they
135 /// are nested inside of any other operations, they are hoisted out of
136 /// them.
137 /// - \c depend, \c device, \c dyn_groupprivate, \c if and \c in_reduction
138 /// clauses are removed from these target functions. Values used to
139 /// initialize other clauses are replaced by placeholders as follows:
140 /// - Values defined by block arguments are replaced by placeholders only
141 /// if they are not attached to the parent function. In that case, they
142 /// are passed unmodified.
143 /// - Pure operations of the LLVM dialect are maintained, and any value
144 /// operands they might have are also replaced by placeholders following
145 /// the same rules.
146 /// - Other values are replaced by new function arguments.
147 /// - \c omp.map.info operations associated to these target regions are
148 /// preserved. These are moved above all \c omp.target and sorted to
149 /// satisfy dependencies among them.
150 /// - \c bounds arguments are removed from \c omp.map.info operations.
151 /// - \c var_ptr and \c var_ptr_ptr arguments of \c omp.map.info are
152 /// replaced by placeholders as described above.
153 /// - Every other operation not located inside of an \c omp.target is
154 /// removed.
155 LogicalResult rewriteHostFunction(LLVM::LLVMFuncOp funcOp) {
156 Region &region = funcOp.getFunctionBody();
157 LLVM::LLVMFunctionType functionType = funcOp.getFunctionType();
158
159 // Collect target operations inside of the function.
160 llvm::SmallVector<omp::TargetOp> targetOps;
161 region.walk<WalkOrder::PreOrder>([&](Operation *op) {
162 // Skip the inside of omp.target regions, since these contain device code.
163 if (auto targetOp = dyn_cast<omp::TargetOp>(op)) {
164 targetOps.push_back(targetOp);
165 return WalkResult::skip();
166 }
167
168 // Replace omp.target_data entry block argument uses with the value used
169 // to initialize the associated omp.map.info operation. This way,
170 // references are still valid once the omp.target operation has been
171 // extracted out of the omp.target_data region.
172 if (auto targetDataOp = dyn_cast<omp::TargetDataOp>(op)) {
173 llvm::SmallVector<std::pair<Value, BlockArgument>> argPairs;
174 cast<omp::BlockArgOpenMPOpInterface>(*targetDataOp)
175 .getBlockArgsPairs(argPairs);
176 for (auto [operand, blockArg] : argPairs) {
177 auto mapInfo = cast<omp::MapInfoOp>(operand.getDefiningOp());
178 blockArg.replaceAllUsesWith(mapInfo.getVarPtr());
179 }
180 }
181 return WalkResult::advance();
182 });
183
184 // Make a temporary clone of the parent function with an empty region,
185 // and update all references to entry block arguments to those of the new
186 // region. Users of these arguments will later either be moved to the new
187 // region or deleted when the original region is replaced by the new.
188 OpBuilder builder(&getContext());
189 builder.setInsertionPointAfter(funcOp);
190 Operation *newFuncOp = builder.cloneWithoutRegions(funcOp);
191 Block &block = newFuncOp->getRegion(0).emplaceBlock();
192
193 llvm::SmallVector<Location> locs;
194 locs.reserve(region.getNumArguments());
195 llvm::transform(region.getArguments(), std::back_inserter(locs),
196 [](const BlockArgument &arg) { return arg.getLoc(); });
197 block.addArguments(region.getArgumentTypes(), locs);
198
199 for (auto [oldArg, newArg] :
200 llvm::zip_equal(region.getArguments(), block.getArguments()))
201 oldArg.replaceAllUsesWith(newArg);
202
203 // Collect omp.map.info ops while satisfying interdependencies and remove
204 // operands that aren't used by target device codegen.
205 //
206 // This logic must be updated whenever operands to omp.target change.
207 llvm::SetVector<Value> rewriteValues;
208 llvm::SetVector<omp::MapInfoOp> mapInfos;
209 for (omp::TargetOp targetOp : targetOps) {
210 assert(targetOp.getHostEvalVars().empty() &&
211 "unexpected host_eval in target device module");
212
213 // Variables unused by the device.
214 targetOp.getDependVarsMutable().clear();
215 targetOp.setDependKindsAttr(nullptr);
216 targetOp.getDependIteratedMutable().clear();
217 targetOp.setDependIteratedKindsAttr(nullptr);
218 targetOp.getDeviceMutable().clear();
219 targetOp.getDynGroupprivateSizeMutable().clear();
220 targetOp.setDynGroupprivateAccessGroupAttr(nullptr);
221 targetOp.setDynGroupprivateFallbackAttr(nullptr);
222 targetOp.getIfExprMutable().clear();
223 targetOp.getInReductionVarsMutable().clear();
224 targetOp.setInReductionByrefAttr(nullptr);
225 targetOp.setInReductionSymsAttr(nullptr);
226
227 // TODO: Clear some of these operands rather than rewriting them,
228 // depending on whether they are needed by device codegen once support for
229 // them is fully implemented.
230 for (Value allocVar : targetOp.getAllocateVars())
231 collectRewrite(allocVar, rewriteValues);
232 for (Value allocVar : targetOp.getAllocatorVars())
233 collectRewrite(allocVar, rewriteValues);
234 for (Value isDevPtr : targetOp.getIsDevicePtrVars())
235 collectRewrite(isDevPtr, rewriteValues);
236 for (Value mapVar : targetOp.getHasDeviceAddrVars())
237 collectRewrite(cast<omp::MapInfoOp>(mapVar.getDefiningOp()), mapInfos);
238 for (Value mapVar : targetOp.getMapVars())
239 collectRewrite(cast<omp::MapInfoOp>(mapVar.getDefiningOp()), mapInfos);
240 for (Value privateVar : targetOp.getPrivateVars())
241 collectRewrite(privateVar, rewriteValues);
242 for (Value threadLimit : targetOp.getThreadLimitVars())
243 collectRewrite(threadLimit, rewriteValues);
244 }
245
246 // Move omp.map.info ops to the new block and collect dependencies.
247 for (omp::MapInfoOp mapOp : mapInfos) {
248 collectRewrite(mapOp.getVarPtr(), rewriteValues);
249
250 if (Value varPtrPtr = mapOp.getVarPtrPtr())
251 collectRewrite(varPtrPtr, rewriteValues);
252
253 // Bounds are not used during target device codegen.
254 mapOp.getBoundsMutable().clear();
255 mapOp->moveBefore(&block, block.end());
256 }
257
258 builder.setInsertionPointToStart(&block);
259
260 // We don't actually need the proper initialization for all operands, but
261 // rather just to maintain the basic form of omp.target operations. We
262 // create new function arguments as placeholders for rewritten values.
263 llvm::SmallVector<Type> newFnArgTypes(functionType.getParams());
264 for (Value value : rewriteValues) {
265 Value rewriteValue;
266 Operation *definingOp = value.getDefiningOp();
267 if (definingOp && keepHostOpInDevice(*definingOp)) {
268 rewriteValue = builder.clone(*value.getDefiningOp())->getResult(0);
269 } else {
270 rewriteValue = block.addArgument(value.getType(), value.getLoc());
271 newFnArgTypes.push_back(rewriteValue.getType());
272 }
273 value.replaceAllUsesWith(rewriteValue);
274 }
275
276 // Move target operations to the end of the new block.
277 for (omp::TargetOp targetOp : targetOps)
278 targetOp->moveBefore(&block, block.end());
279
280 // Add terminator to the new block.
281 builder.setInsertionPointToEnd(&block);
282 LLVM::ReturnOp::create(builder, funcOp.getLoc(), ValueRange());
283
284 // Replace old region with the new one, now only containing the required
285 // operations, and remove the temporary operation clone.
286 region.takeBody(newFuncOp->getRegion(0));
287 newFuncOp->erase();
288
289 // Update function type after modifying the terminator and argument list.
290 funcOp.setType(LLVM::LLVMFunctionType::get(
291 LLVM::LLVMVoidType::get(&getContext()), newFnArgTypes));
292
293 return success();
294 }
295};
296} // namespace
return success()
static std::optional< omp::DeclareTargetDeviceType > getDeclareTargetDevice(Operation &op)
Provide the device_type of an omp.declare_target attribute, if defined.
static bool keepHostOpInDevice(Operation &op)
Some host operations, like llvm.mlir.addressof and constants, must remain in the device module becaus...
static void collectRewrite(omp::MapInfoOp mapOp, llvm::SetVector< omp::MapInfoOp > &rewrites)
Add an omp.map.info operation and all its members recursively to the output set to be later rewritten...
b getContext())
iterator_range< args_iterator > addArguments(TypeRange types, ArrayRef< Location > locs)
Add one argument to the argument list for each type specified in the list.
Definition Block.cpp:165
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Definition Block.cpp:158
BlockArgListType getArguments()
Definition Block.h:112
iterator end()
Definition Block.h:169
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Definition Operation.h:237
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
void erase()
Remove this operation from its parent block and delete it.
BlockArgListType getArguments()
Definition Region.h:94
Block & emplaceBlock()
Definition Region.h:46
unsigned getNumArguments()
Definition Region.h:136
ValueTypeRange< BlockArgListType > getArgumentTypes()
Returns the argument types of the first block within the region.
Definition Region.cpp:36
void takeBody(Region &other)
Takes body of another region (that region will have no body after this operation completes).
Definition Region.h:253
RetT walk(FnT &&callback)
Walk all nested operations, blocks or regions (including this region), depending on the type of callb...
Definition Region.h:297
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
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool isPure(Operation *op)
Returns true if the given operation is pure, i.e., is speculatable that does not touch memory.