22#define GEN_PASS_DEF_HOSTOPFILTERINGPASS
23#include "mlir/Dialect/OpenMP/Transforms/Passes.h.inc"
49 for (
Value member : mapOp.getMembers())
50 collectRewrite(cast<omp::MapInfoOp>(member.getDefiningOp()), rewrites);
52 rewrites.insert(mapOp);
62 if ((isa<BlockArgument>(value) &&
63 isa<FunctionOpInterface>(
64 cast<BlockArgument>(value).getOwner()->getParentOp())) ||
65 rewrites.contains(value))
73 rewrites.insert(value);
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();
89class HostOpFilteringPass
90 :
public omp::impl::HostOpFilteringPassBase<HostOpFilteringPass> {
92 HostOpFilteringPass() =
default;
94 void runOnOperation()
override {
95 auto op = dyn_cast<omp::OffloadModuleInterface>(getOperation());
96 if (!op || !op.getIsTargetDevice())
99 op->walk<WalkOrder::PreOrder>([&](LLVM::LLVMFuncOp funcOp) {
100 omp::DeclareTargetDeviceType declareType =
102 .value_or(omp::DeclareTargetDeviceType::host);
105 if (funcOp.isExternal() ||
106 declareType != omp::DeclareTargetDeviceType::host)
109 if (
failed(rewriteHostFunction(funcOp))) {
110 funcOp.emitOpError() <<
"could not filter host-only operations";
120 op->walk([&](LLVM::GlobalOp globalOp) {
122 globalOp.setLinkage(LLVM::Linkage::Internal);
155 LogicalResult rewriteHostFunction(LLVM::LLVMFuncOp funcOp) {
156 Region ®ion = funcOp.getFunctionBody();
157 LLVM::LLVMFunctionType functionType = funcOp.getFunctionType();
160 llvm::SmallVector<omp::TargetOp> targetOps;
161 region.
walk<WalkOrder::PreOrder>([&](Operation *op) {
163 if (
auto targetOp = dyn_cast<omp::TargetOp>(op)) {
164 targetOps.push_back(targetOp);
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());
189 builder.setInsertionPointAfter(funcOp);
190 Operation *newFuncOp = builder.cloneWithoutRegions(funcOp);
193 llvm::SmallVector<Location> locs;
195 llvm::transform(region.
getArguments(), std::back_inserter(locs),
196 [](
const BlockArgument &arg) { return arg.getLoc(); });
199 for (
auto [oldArg, newArg] :
201 oldArg.replaceAllUsesWith(newArg);
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");
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);
230 for (Value allocVar : targetOp.getAllocateVars())
232 for (Value allocVar : targetOp.getAllocatorVars())
234 for (Value isDevPtr : targetOp.getIsDevicePtrVars())
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())
242 for (Value threadLimit : targetOp.getThreadLimitVars())
247 for (omp::MapInfoOp mapOp : mapInfos) {
250 if (Value varPtrPtr = mapOp.getVarPtrPtr())
254 mapOp.getBoundsMutable().clear();
255 mapOp->moveBefore(&block, block.
end());
258 builder.setInsertionPointToStart(&block);
263 llvm::SmallVector<Type> newFnArgTypes(functionType.getParams());
264 for (Value value : rewriteValues) {
266 Operation *definingOp = value.getDefiningOp();
268 rewriteValue = builder.clone(*value.getDefiningOp())->getResult(0);
270 rewriteValue = block.
addArgument(value.getType(), value.getLoc());
271 newFnArgTypes.push_back(rewriteValue.
getType());
273 value.replaceAllUsesWith(rewriteValue);
277 for (omp::TargetOp targetOp : targetOps)
278 targetOp->moveBefore(&block, block.
end());
281 builder.setInsertionPointToEnd(&block);
282 LLVM::ReturnOp::create(builder, funcOp.getLoc(),
ValueRange());
290 funcOp.setType(LLVM::LLVMFunctionType::get(
291 LLVM::LLVMVoidType::get(&
getContext()), newFnArgTypes));
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...
iterator_range< args_iterator > addArguments(TypeRange types, ArrayRef< Location > locs)
Add one argument to the argument list for each type specified in the list.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
BlockArgListType getArguments()
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
Operation is the basic unit of execution within MLIR.
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
operand_range getOperands()
Returns an iterator on the underlying Value's.
MLIRContext * getContext()
Return the context this operation is associated with.
void erase()
Remove this operation from its parent block and delete it.
BlockArgListType getArguments()
unsigned getNumArguments()
ValueTypeRange< BlockArgListType > getArgumentTypes()
Returns the argument types of the first block within the region.
void takeBody(Region &other)
Takes body of another region (that region will have no body after this operation completes).
RetT walk(FnT &&callback)
Walk all nested operations, blocks or regions (including this region), depending on the type of callb...
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.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static WalkResult advance()
static WalkResult interrupt()
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.