[flang-commits] [flang] [Flang][OpenMP] Optimize target updates of derived-type scalars (PR #219488)
via flang-commits
flang-commits at lists.llvm.org
Fri Aug 28 09:55:44 PDT 2026
================
@@ -4300,6 +4302,173 @@ static mlir::omp::TargetDataOp genTargetDataOp(
return targetDataOp;
}
+struct TargetUpdateKernelEntry {
+ mlir::omp::MapInfoOp mapInfo;
+ mlir::Value hostPtr;
+ mlir::Type componentType;
+};
+
+static bool hasAMDGCNTarget(mlir::ModuleOp module) {
+ auto offloadModule = llvm::cast<mlir::omp::OffloadModuleInterface>(*module);
+ if (offloadModule.getIsTargetDevice())
+ return fir::getTargetTriple(module).isAMDGCN();
+ return llvm::any_of(
+ offloadModule.getTargetTriples(), [](mlir::Attribute attr) {
+ auto tripleAttr = llvm::dyn_cast<mlir::StringAttr>(attr);
+ return tripleAttr && llvm::Triple(tripleAttr.getValue()).isAMDGCN();
+ });
+}
+
+static std::optional<TargetUpdateKernelEntry>
+getTargetUpdateKernelEntry(mlir::Value mapVar) {
+ auto mapInfo = mapVar.getDefiningOp<mlir::omp::MapInfoOp>();
+ if (!mapInfo)
+ return std::nullopt;
+
+ // Keep the fast path to plain synchronous H2D motion. In particular, do not
+ // silently weaken `present` motion modifiers.
+ if (mapInfo.getMapType() != mlir::omp::ClauseMapFlags::to ||
+ mapInfo.getVarPtrPtr() || !mapInfo.getMembers().empty() ||
+ !mapInfo.getBounds().empty() || mapInfo.getMapperId())
+ return std::nullopt;
+
+ mlir::Value hostPtr = mapInfo.getVarPtr();
+ auto designate = hostPtr.getDefiningOp<hlfir::DesignateOp>();
+ if (!designate || !designate.getComponent() ||
+ designate.getComponentShape() || !designate.getIndices().empty() ||
+ !designate.getSubstring().empty() || designate.getComplexPart() ||
+ designate.getShape() || !designate.getTypeparams().empty())
+ return std::nullopt;
+
+ mlir::Type baseType = fir::unwrapRefType(designate.getMemref().getType());
+ auto recordType = mlir::dyn_cast<fir::RecordType>(baseType);
+ if (!recordType || recordType.getNumLenParams() != 0)
+ return std::nullopt;
+
+ llvm::StringRef component = designate.getComponent()->getValue();
+ mlir::Type componentType = recordType.getType(component);
+ if (!componentType || !fir::isa_trivial(componentType))
+ return std::nullopt;
+
+ return TargetUpdateKernelEntry{mapInfo, hostPtr, componentType};
+}
+
+static mlir::omp::TargetOp
+genTargetUpdateKernel(lower::AbstractConverter &converter, mlir::Location loc,
+ llvm::ArrayRef<TargetUpdateKernelEntry> entries) {
+ fir::FirOpBuilder &builder = converter.getFirOpBuilder();
+ mlir::omp::TargetExtOperands targetClauseOps;
+ targetClauseOps.kernelType = mlir::omp::TargetExecModeAttr::get(
+ builder.getContext(), mlir::omp::TargetExecMode::generic);
+
+ llvm::SmallVector<mlir::Value> destinationMaps;
+ destinationMaps.reserve(entries.size());
+
+ llvm::SmallVector<mlir::Type> sourceTypes;
+ llvm::transform(
+ entries, std::back_inserter(sourceTypes),
+ [](const TargetUpdateKernelEntry &entry) { return entry.componentType; });
+ mlir::TupleType sourceType =
+ mlir::TupleType::get(builder.getContext(), sourceTypes);
+ mlir::Value sourcePack = builder.createTemporary(loc, sourceType);
+
+ for (auto indexedEntry : llvm::enumerate(entries)) {
----------------
agozillon wrote:
Nit: might be able to use the [i, &entry] syntax to save you the accesses in the loop
https://github.com/llvm/llvm-project/pull/219488
More information about the flang-commits
mailing list