[flang-commits] [flang] [Flang][OpenMP] Optimize target updates of derived-type scalars (PR #219488)
Akash Banerjee via flang-commits
flang-commits at lists.llvm.org
Fri Aug 28 07:28:49 PDT 2026
https://github.com/TIFitis created https://github.com/llvm/llvm-project/pull/219488
Pack AMDGPU derived-type scalar updates into a single transfer and generated target region, avoiding one costly runtime transfer per component. Preserve existing lowering for unsupported clauses, pointer components, and non-AMDGPU targets.
>From 43775aabab5bb23d0057feca92251fb50377ca91 Mon Sep 17 00:00:00 2001
From: Akash Banerjee <Akash.Banerjee at amd.com>
Date: Tue, 25 Aug 2026 17:21:31 +0100
Subject: [PATCH] [Flang][OpenMP] Optimize target updates of derived-type
scalars
Pack AMDGPU derived-type scalar updates into a single transfer and generated target region, avoiding one costly runtime transfer per component. Preserve existing lowering for unsupported clauses, pointer components, and non-AMDGPU targets.
---
flang/lib/Lower/OpenMP/OpenMP.cpp | 191 +++++++++++++++++-
.../OpenMP/target-update-derived-type.f90 | 114 +++++++++++
2 files changed, 303 insertions(+), 2 deletions(-)
create mode 100644 flang/test/Lower/OpenMP/target-update-derived-type.f90
diff --git a/flang/lib/Lower/OpenMP/OpenMP.cpp b/flang/lib/Lower/OpenMP/OpenMP.cpp
index 7503d33c8df38..684db4340e5d4 100644
--- a/flang/lib/Lower/OpenMP/OpenMP.cpp
+++ b/flang/lib/Lower/OpenMP/OpenMP.cpp
@@ -39,6 +39,7 @@
#include "flang/Optimizer/Builder/Todo.h"
#include "flang/Optimizer/Dialect/FIROpsSupport.h"
#include "flang/Optimizer/Dialect/FIRType.h"
+#include "flang/Optimizer/Dialect/Support/FIRContext.h"
#include "flang/Optimizer/HLFIR/HLFIROps.h"
#include "flang/Optimizer/Support/InternalNames.h"
#include "flang/Parser/openmp-utils.h"
@@ -65,6 +66,7 @@
#include "llvm/ADT/SmallSet.h"
#include "llvm/ADT/StringSwitch.h"
#include "llvm/Frontend/OpenMP/OMP.h"
+#include "llvm/TargetParser/Triple.h"
#include <atomic>
using namespace Fortran::lower::omp;
@@ -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)) {
+ std::size_t i = indexedEntry.index();
+ const TargetUpdateKernelEntry &entry = indexedEntry.value();
+ mlir::Value sourceValue = fir::LoadOp::create(builder, loc, entry.hostPtr);
+ mlir::Value index =
+ builder.createIntegerConstant(loc, builder.getI32Type(), i);
+ mlir::Value sourceAddr = fir::CoordinateOp::create(
+ builder, loc, builder.getRefType(entry.componentType), sourcePack,
+ index);
+ fir::StoreOp::create(builder, loc, sourceValue, sourceAddr);
+
+ mlir::Value destinationMap = createMapInfoOp(
+ builder, loc, entry.hostPtr, /*varPtrPtr=*/mlir::Value{},
+ /*name=*/"", /*bounds=*/{}, /*members=*/{},
+ /*membersIndex=*/mlir::ArrayAttr{}, mlir::omp::ClauseMapFlags::storage,
+ mlir::omp::VariableCaptureKind::ByRef, entry.hostPtr.getType());
+ destinationMaps.push_back(destinationMap);
+ }
+
+ mlir::Value sourceMap = createMapInfoOp(
+ builder, loc, sourcePack, /*varPtrPtr=*/mlir::Value{},
+ ".omp.target.update.source", /*bounds=*/{}, /*members=*/{},
+ /*membersIndex=*/mlir::ArrayAttr{}, mlir::omp::ClauseMapFlags::to,
+ mlir::omp::VariableCaptureKind::ByRef, sourcePack.getType());
+ targetClauseOps.mapVars.push_back(sourceMap);
+ targetClauseOps.mapVars.append(destinationMaps);
+
+ auto targetOp = mlir::omp::TargetOp::create(builder, loc, targetClauseOps);
+ llvm::SmallVector<mlir::Value> mapBaseValues;
+ extractMappedBaseValues(targetClauseOps.mapVars, mapBaseValues);
+ ObjectEntryBlockArgs args;
+ args.map.vars = mapBaseValues;
+ genEntryBlock(builder, args.asEntryBlockArgs(), targetOp.getRegion());
+
+ auto argIface = llvm::cast<mlir::omp::BlockArgOpenMPOpInterface>(*targetOp);
+ llvm::ArrayRef<mlir::BlockArgument> mapBlockArgs = argIface.getMapBlockArgs();
+ assert(mapBlockArgs.size() == entries.size() + 1 &&
+ "expected source and destination map arguments");
+ builder.setInsertionPointToEnd(&targetOp.getRegion().front());
+ for (unsigned i = 0; i < entries.size(); ++i) {
+ mlir::Value index =
+ builder.createIntegerConstant(loc, builder.getI32Type(), i);
+ mlir::Value sourceAddr = fir::CoordinateOp::create(
+ builder, loc, builder.getRefType(entries[i].componentType),
+ mapBlockArgs.front(), index);
+ mlir::Value sourceValue = fir::LoadOp::create(builder, loc, sourceAddr);
+ fir::StoreOp::create(builder, loc, sourceValue, mapBlockArgs[i + 1]);
+ }
+ mlir::omp::TerminatorOp::create(builder, loc);
+ builder.setInsertionPointAfter(targetOp);
+ return targetOp;
+}
+
+static mlir::Operation *tryGenTargetUpdateKernel(
+ lower::AbstractConverter &converter, mlir::Location loc,
+ mlir::omp::TargetEnterExitUpdateDataOperands &clauseOps) {
+ // Updating several small, discontiguous fields issues one device transfer
+ // for every map entry. Pack their host values and use one target region so
+ // that the runtime performs one H2D transfer followed by the scalar stores.
+ // This addresses the AMDGPU runtime transfer cost and is only enabled when
+ // an AMDGPU image will actually be emitted.
+ if (!hasAMDGCNTarget(converter.getModuleOp()) || clauseOps.mapVars.empty() ||
+ !clauseOps.dependVars.empty() || !clauseOps.dependIterated.empty() ||
+ !clauseOps.mapIterated.empty() || clauseOps.nowait || clauseOps.device)
+ return nullptr;
+
+ llvm::SmallVector<TargetUpdateKernelEntry> entries;
+ entries.reserve(clauseOps.mapVars.size());
+ for (mlir::Value mapVar : clauseOps.mapVars) {
+ std::optional<TargetUpdateKernelEntry> entry =
+ getTargetUpdateKernelEntry(mapVar);
+ if (!entry)
+ return nullptr;
+ entries.push_back(*entry);
+ }
+
+ fir::FirOpBuilder &builder = converter.getFirOpBuilder();
+ mlir::Operation *firstGenerated = nullptr;
+
+ if (mlir::Value ifExpr = clauseOps.ifExpr) {
+ auto ifOp = fir::IfOp::create(builder, loc, ifExpr,
+ /*withElseRegion=*/false);
+ firstGenerated = ifOp;
+ builder.setInsertionPoint(ifOp.getThenRegion().front().getTerminator());
+ genTargetUpdateKernel(converter, loc, entries);
+ builder.setInsertionPointAfter(ifOp);
+ } else {
+ firstGenerated = genTargetUpdateKernel(converter, loc, entries);
+ }
+
+ for (TargetUpdateKernelEntry &entry : entries)
+ if (entry.mapInfo->use_empty())
+ entry.mapInfo.erase();
+
+ return firstGenerated;
+}
+
template <typename OpTy>
static OpTy genTargetEnterExitUpdateDataOp(
lower::AbstractConverter &converter, lower::SymMap &symTable,
@@ -4327,6 +4496,24 @@ static OpTy genTargetEnterExitUpdateDataOp(
return OpTy::create(firOpBuilder, loc, clauseOps);
}
+static mlir::Operation *
+genTargetUpdateDataOp(lower::AbstractConverter &converter,
+ lower::SymMap &symTable, lower::StatementContext &stmtCtx,
+ semantics::SemanticsContext &semaCtx, mlir::Location loc,
+ const ConstructQueue &queue,
+ ConstructQueue::const_iterator item) {
+ fir::FirOpBuilder &firOpBuilder = converter.getFirOpBuilder();
+ mlir::omp::TargetEnterExitUpdateDataOperands clauseOps;
+ genTargetEnterExitUpdateDataClauses(
+ converter, semaCtx, symTable, stmtCtx, item->clauses, loc,
+ llvm::omp::Directive::OMPD_target_update, clauseOps);
+
+ if (mlir::Operation *op = tryGenTargetUpdateKernel(converter, loc, clauseOps))
+ return op;
+
+ return mlir::omp::TargetUpdateOp::create(firOpBuilder, loc, clauseOps);
+}
+
static mlir::omp::TaskOp
genTaskOp(lower::AbstractConverter &converter, lower::SymMap &symTable,
lower::StatementContext &stmtCtx,
@@ -5501,8 +5688,8 @@ static void genOMPDispatch(lower::AbstractConverter &converter,
converter, symTable, stmtCtx, semaCtx, loc, queue, item);
break;
case llvm::omp::Directive::OMPD_target_update:
- newOp = genTargetEnterExitUpdateDataOp<mlir::omp::TargetUpdateOp>(
- converter, symTable, stmtCtx, semaCtx, loc, queue, item);
+ newOp = genTargetUpdateDataOp(converter, symTable, stmtCtx, semaCtx, loc,
+ queue, item);
break;
case llvm::omp::Directive::OMPD_task:
newOp = genTaskOp(converter, symTable, stmtCtx, semaCtx, eval, loc, queue,
diff --git a/flang/test/Lower/OpenMP/target-update-derived-type.f90 b/flang/test/Lower/OpenMP/target-update-derived-type.f90
new file mode 100644
index 0000000000000..4552e14ffb8ee
--- /dev/null
+++ b/flang/test/Lower/OpenMP/target-update-derived-type.f90
@@ -0,0 +1,114 @@
+! RUN: %flang_fc1 -emit-hlfir -fopenmp -fopenmp-targets=amdgcn-amd-amdhsa %s -o - | FileCheck %s
+! RUN: %flang_fc1 -emit-hlfir -fopenmp %s -o - | FileCheck %s --check-prefix=HOST
+! RUN: %flang_fc1 -emit-hlfir -fopenmp -fopenmp-targets=nvptx64-nvidia-cuda %s -o - | FileCheck %s --check-prefix=NONAMD
+! RUN: %flang_fc1 -triple amdgcn-amd-amdhsa -emit-hlfir -fopenmp -fopenmp-is-target-device %s -o - | FileCheck %s --check-prefix=DEVICE
+
+module target_update_derived_type
+ type :: wavefun
+ real(8) :: ferwe
+ real(8) :: aux
+ complex(8) :: celen
+ integer :: nb
+ integer :: isp
+ logical :: ldo
+ integer, pointer :: ptr
+ end type
+contains
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_with_if(
+! DEVICE-LABEL: func.func @_QMtarget_update_derived_typePupdate_with_if(
+! DEVICE: omp.target kernel_type(generic)
+! DEVICE-NOT: omp.target_update
+subroutine update_with_if(w, enabled)
+ type(wavefun) :: w
+ logical :: enabled
+
+ ! CHECK: %[[SOURCE:.*]] = fir.alloca tuple<f64, complex<f64>, i32, i32, !fir.logical<4>>
+ ! CHECK: %[[COND:.*]] = fir.convert %{{.*}} : (!fir.logical<4>) -> i1
+ ! CHECK: %[[FERWE:.*]] = hlfir.designate %{{.*}}{"ferwe"}
+ ! CHECK: %[[CELEN:.*]] = hlfir.designate %{{.*}}{"celen"}
+ ! CHECK: fir.if %[[COND]] {
+ ! CHECK: fir.store {{.*}} to {{.*}} : !fir.ref<f64>
+ ! CHECK: %[[FERWE_MAP:.*]] = omp.map.info var_ptr(%[[FERWE]] : !fir.ref<f64>, f64) map_clauses(storage) capture(ByRef)
+ ! CHECK: fir.store {{.*}} to {{.*}} : !fir.ref<complex<f64>>
+ ! CHECK: %[[CELEN_MAP:.*]] = omp.map.info var_ptr(%[[CELEN]] : !fir.ref<complex<f64>>, complex<f64>) map_clauses(storage) capture(ByRef)
+ ! CHECK: %[[SOURCE_MAP:.*]] = omp.map.info var_ptr(%[[SOURCE]] {{.*}}) map_clauses(to) capture(ByRef) name(".omp.target.update.source")
+ ! CHECK: omp.target kernel_type(generic) map_entries(%[[SOURCE_MAP]] -> [[SOURCE_ARG:%[^, ]+]], %[[FERWE_MAP]] -> [[FERWE_ARG:%[^, ]+]], %[[CELEN_MAP]] -> [[CELEN_ARG:%[^, ]+]]
+ ! CHECK: %[[FERWE_SOURCE:.*]] = fir.coordinate_of [[SOURCE_ARG]], {{.*}} -> !fir.ref<f64>
+ ! CHECK: %[[FERWE_VALUE:.*]] = fir.load %[[FERWE_SOURCE]] : !fir.ref<f64>
+ ! CHECK: fir.store %[[FERWE_VALUE]] to [[FERWE_ARG]] : !fir.ref<f64>
+ ! CHECK: %[[CELEN_SOURCE:.*]] = fir.coordinate_of [[SOURCE_ARG]], {{.*}} -> !fir.ref<complex<f64>>
+ ! CHECK: %[[CELEN_VALUE:.*]] = fir.load %[[CELEN_SOURCE]] : !fir.ref<complex<f64>>
+ ! CHECK: fir.store %[[CELEN_VALUE]] to [[CELEN_ARG]] : !fir.ref<complex<f64>>
+ ! CHECK-NOT: omp.target_update
+ ! CHECK: return
+ ! HOST: omp.target_update
+ ! NONAMD: omp.target_update
+ !$omp target update to(w%ferwe, w%celen, w%nb, w%isp, w%ldo) if(enabled)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_without_if(
+! DEVICE-LABEL: func.func @_QMtarget_update_derived_typePupdate_without_if(
+subroutine update_without_if(w)
+ type(wavefun) :: w
+
+ ! CHECK: %[[SOURCE:.*]] = fir.alloca tuple<complex<f64>, i32, i32, !fir.logical<4>>
+ ! CHECK: %[[CELEN:.*]] = hlfir.designate %{{.*}}{"celen"}
+ ! CHECK: %[[CELEN_MAP:.*]] = omp.map.info var_ptr(%[[CELEN]] : !fir.ref<complex<f64>>, complex<f64>) map_clauses(storage) capture(ByRef)
+ ! CHECK: %[[SOURCE_MAP:.*]] = omp.map.info var_ptr(%[[SOURCE]] {{.*}}) map_clauses(to) capture(ByRef) name(".omp.target.update.source")
+ ! CHECK: omp.target kernel_type(generic) map_entries(%[[SOURCE_MAP]] -> [[SOURCE_ARG:%[^, ]+]], %[[CELEN_MAP]] -> [[CELEN_ARG:%[^, ]+]]
+ ! CHECK: %[[CELEN_SOURCE:.*]] = fir.coordinate_of [[SOURCE_ARG]], {{.*}} -> !fir.ref<complex<f64>>
+ ! CHECK: %[[CELEN_VALUE:.*]] = fir.load %[[CELEN_SOURCE]] : !fir.ref<complex<f64>>
+ ! CHECK: fir.store %[[CELEN_VALUE]] to [[CELEN_ARG]] : !fir.ref<complex<f64>>
+ ! CHECK-NOT: omp.target_update
+ ! CHECK: return
+ !$omp target update to(w%celen, w%nb, w%isp, w%ldo)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_pointer(
+subroutine update_pointer(w)
+ type(wavefun) :: w
+
+ ! CHECK: omp.map.info {{.*}} map_clauses(to)
+ ! CHECK: omp.target_update map_entries(
+ !$omp target update to(w%ptr)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_device(
+subroutine update_device(w)
+ type(wavefun) :: w
+
+ ! CHECK: %[[MAP:.*]] = omp.map.info {{.*}} map_clauses(to)
+ ! CHECK: omp.target_update device({{.*}}) map_entries(%[[MAP]]
+ !$omp target update to(w%ferwe) device(0)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_array_element(
+subroutine update_array_element(w)
+ type(wavefun) :: w(2)
+
+ ! CHECK: fir.alloca tuple<f64, i32>
+ ! CHECK: omp.target kernel_type(generic)
+ ! CHECK-NOT: omp.target_update
+ !$omp target update to(w(2)%ferwe, w(2)%nb)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_from(
+subroutine update_from(w)
+ type(wavefun) :: w
+
+ ! CHECK: %[[MAP:.*]] = omp.map.info {{.*}} map_clauses(from)
+ ! CHECK: omp.target_update map_entries(%[[MAP]]
+ !$omp target update from(w%ferwe)
+end subroutine
+
+! CHECK-LABEL: func.func @_QMtarget_update_derived_typePupdate_nowait(
+subroutine update_nowait(w)
+ type(wavefun) :: w
+
+ ! CHECK: %[[MAP:.*]] = omp.map.info {{.*}} map_clauses(to)
+ ! CHECK: omp.target_update map_entries(%[[MAP]]{{.*}}) nowait
+ !$omp target update to(w%ferwe) nowait
+end subroutine
+
+end module
More information about the flang-commits
mailing list