[flang-commits] [flang] [Flang][OpenMP] Fix nested OpenMP mapper composition (PR #215639)

Akash Banerjee via flang-commits flang-commits at lists.llvm.org
Wed Aug 12 09:55:02 PDT 2026


https://github.com/TIFitis updated https://github.com/llvm/llvm-project/pull/215639

>From 5bbc3ae419aa44e4655b4b6b44883e971872aae3 Mon Sep 17 00:00:00 2001
From: Akash Banerjee <Akash.Banerjee at amd.com>
Date: Tue, 11 Aug 2026 19:18:37 +0100
Subject: [PATCH 1/3] [Flang][OpenMP] Fix nested OpenMP mapper composition

Fixes Flang lowering for nested OpenMP declare mappers by attaching the nested default mapper to mapped derived-type components and preventing redundant implicit member maps already covered by that mapper.
---
 flang/lib/Lower/OpenMP/Utils.cpp              |  6 +-
 .../Optimizer/OpenMP/MapInfoFinalization.cpp  | 87 ++++++++++++++++---
 .../Lower/OpenMP/nested-default-mapper.f90    | 37 ++++++++
 3 files changed, 113 insertions(+), 17 deletions(-)
 create mode 100644 flang/test/Lower/OpenMP/nested-default-mapper.f90

diff --git a/flang/lib/Lower/OpenMP/Utils.cpp b/flang/lib/Lower/OpenMP/Utils.cpp
index 8f57f00d59c58..88ddf504bdd98 100644
--- a/flang/lib/Lower/OpenMP/Utils.cpp
+++ b/flang/lib/Lower/OpenMP/Utils.cpp
@@ -1089,9 +1089,9 @@ static std::string
 getDefaultMapperID(Fortran::lower::AbstractConverter &converter,
                    fir::FirOpBuilder &firOpBuilder,
                    const semantics::DerivedTypeSpec *typeSpec) {
-  if (mlir::isa<mlir::omp::DeclareMapperOp>(
-          firOpBuilder.getRegion().getParentOp()) ||
-      !typeSpec)
+  // Nested derived-type components inside a declare mapper may use their own
+  // default mapper. Only suppress the mapper currently being built.
+  if (!typeSpec)
     return {};
 
   std::string mapperIdName =
diff --git a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
index 949da8f20cbbe..2a77657b90adb 100644
--- a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
+++ b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
@@ -138,6 +138,75 @@ class MapInfoFinalizationPass
     return findMemberByIndexPath(op, indexPath) != nullptr;
   }
 
+  static bool mapperCoversIndexPath(
+      mlir::Operation *symbolTableAnchor, mlir::FlatSymbolRefAttr mapperId,
+      llvm::ArrayRef<int64_t> indexPath,
+      llvm::SmallPtrSetImpl<mlir::Operation *> &visitedMappers) {
+    mlir::omp::DeclareMapperOp symbol =
+        mlir::SymbolTable::lookupNearestSymbolFrom<mlir::omp::DeclareMapperOp>(
+            symbolTableAnchor, mapperId);
+    if (!symbol || !visitedMappers.insert(symbol.getOperation()).second)
+      return false;
+
+    mlir::omp::DeclareMapperInfoOp mapperInfo = symbol.getDeclareMapperInfo();
+    if (!mapperInfo)
+      return false;
+
+    return llvm::any_of(mapperInfo.getMapVars(), [&](mlir::Value v) {
+      mlir::omp::MapInfoOp map =
+          mlir::dyn_cast_if_present<mlir::omp::MapInfoOp>(v.getDefiningOp());
+      return map && !map.getMembers().empty() &&
+             map.getMembersIndexAttr() &&
+             mapInfoCoversIndexPath(map, indexPath, visitedMappers);
+    });
+  }
+
+  static bool mapInfoCoversIndexPath(
+      mlir::omp::MapInfoOp map, llvm::ArrayRef<int64_t> indexPath,
+      llvm::SmallPtrSetImpl<mlir::Operation *> &visitedMappers) {
+    if (mappedIndexPathExists(map, indexPath))
+      return true;
+
+    mlir::ArrayAttr memberIndices = map.getMembersIndexAttr();
+    if (!memberIndices)
+      return false;
+
+    for (auto [memberIdx, memberIndexAttr] : llvm::enumerate(memberIndices)) {
+      auto memberIndexPath = mlir::cast<mlir::ArrayAttr>(memberIndexAttr);
+      if (memberIndexPath.size() >= indexPath.size())
+        continue;
+
+      bool isPrefix = true;
+      for (auto [idx, attr] : llvm::enumerate(memberIndexPath)) {
+        if (mlir::cast<mlir::IntegerAttr>(attr).getInt() != indexPath[idx]) {
+          isPrefix = false;
+          break;
+        }
+      }
+      if (!isPrefix)
+        continue;
+
+      mlir::omp::MapInfoOp memberMap =
+          mlir::dyn_cast_if_present<mlir::omp::MapInfoOp>(
+              map.getMembers()[memberIdx].getDefiningOp());
+      if (!memberMap)
+        continue;
+
+      llvm::ArrayRef<int64_t> nestedIndexPath =
+          indexPath.drop_front(memberIndexPath.size());
+      if (!memberMap.getMembers().empty() &&
+          mapInfoCoversIndexPath(memberMap, nestedIndexPath, visitedMappers))
+        return true;
+
+      if (memberMap.getMapperIdAttr() &&
+          mapperCoversIndexPath(memberMap, memberMap.getMapperIdAttr(),
+                                nestedIndexPath, visitedMappers))
+        return true;
+    }
+
+    return false;
+  }
+
   /// Get the map type of the nearest explicitly mapped parent for a member.
   /// "Explicitly mapped" means the map type does NOT have the implicit flag.
   ///
@@ -205,20 +274,10 @@ class MapInfoFinalizationPass
       return;
 
     if (op.getMapperId()) {
-      mlir::omp::DeclareMapperOp symbol =
-          mlir::SymbolTable::lookupNearestSymbolFrom<
-              mlir::omp::DeclareMapperOp>(op, op.getMapperIdAttr());
-      assert(symbol && "missing symbol for declare mapper identifier");
-      mlir::omp::DeclareMapperInfoOp mapperInfo = symbol.getDeclareMapperInfo();
-      // TODO: Probably a way to cache these keys in someway so we don't
-      // constantly go through the process of rebuilding them on every check, to
-      // save some cycles, but it can wait for a subsequent patch.
-      for (auto v : mapperInfo.getMapVars()) {
-        mlir::omp::MapInfoOp map =
-            mlir::cast<mlir::omp::MapInfoOp>(v.getDefiningOp());
-        if (!map.getMembers().empty() && mappedIndexPathExists(map, indexPath))
-          return;
-      }
+      llvm::SmallPtrSet<mlir::Operation *, 4> visitedMappers;
+      if (mapperCoversIndexPath(op, op.getMapperIdAttr(), indexPath,
+                                visitedMappers))
+        return;
     }
 
     builder.setInsertionPoint(op);
diff --git a/flang/test/Lower/OpenMP/nested-default-mapper.f90 b/flang/test/Lower/OpenMP/nested-default-mapper.f90
new file mode 100644
index 0000000000000..1dc0a5bc35aad
--- /dev/null
+++ b/flang/test/Lower/OpenMP/nested-default-mapper.f90
@@ -0,0 +1,37 @@
+! RUN: %flang_fc1 -emit-hlfir -fopenmp -fopenmp-version=50 %s -o - | FileCheck %s
+
+program main
+  implicit none
+
+  type nested_t
+    integer, allocatable :: y(:)
+  end type nested_t
+
+  !$omp declare mapper(nested_t :: n) map(n%y)
+
+  type typ_t
+    integer, allocatable :: x(:)
+    type(nested_t) :: nested
+  end type typ_t
+
+  !$omp declare mapper(typ_t :: t) map(t%x, t%nested)
+
+  type(typ_t) :: typ
+
+  allocate(typ%x(3), source=1)
+  allocate(typ%nested%y(3), source=42)
+
+  !$omp target map(tofrom: typ)
+    typ%x(1) = 999
+    typ%nested%y(1) = -555
+  !$omp end target
+end program main
+
+! CHECK-LABEL: omp.declare_mapper @_QQFtyp_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} map_clauses(tofrom) capture(ByRef) mapper(@_QQFnested_t_omp_default_mapper) -> {{.*}} {name = "t%nested"}
+! CHECK-LABEL: omp.declare_mapper @_QQFnested_t_omp_default_mapper
+
+! CHECK-LABEL: func.func @_QQmain
+! CHECK-NOT: implicit_map
+! CHECK: %[[TYP_MAP:.*]] = omp.map.info {{.*}} mapper(@_QQFtyp_t_omp_default_mapper){{.*}} {name = "typ"}
+! CHECK-NEXT: omp.target kernel_type(generic) map_entries(%[[TYP_MAP]] -> %{{[^,]*}} : {{.*}}) {

>From de31420043a170d39fb8e65c2f266043bda7cd68 Mon Sep 17 00:00:00 2001
From: Akash Banerjee <Akash.Banerjee at amd.com>
Date: Wed, 12 Aug 2026 17:26:05 +0100
Subject: [PATCH 2/3] Fix recursive mapper composition.

---
 .../Optimizer/OpenMP/MapInfoFinalization.cpp  | 54 +++++++++++------
 .../Lower/OpenMP/nested-default-mapper.f90    | 60 +++++++++++++++++++
 2 files changed, 96 insertions(+), 18 deletions(-)

diff --git a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
index 2a77657b90adb..0340e32d0438f 100644
--- a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
+++ b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
@@ -41,7 +41,6 @@
 #include "mlir/Pass/Pass.h"
 #include "mlir/Support/LLVM.h"
 #include "llvm/ADT/BitmaskEnum.h"
-#include "llvm/ADT/SmallPtrSet.h"
 #include "llvm/ADT/StringSet.h"
 #include "llvm/Support/raw_ostream.h"
 #include <algorithm>
@@ -67,6 +66,10 @@ class MapInfoFinalizationPass
     size_t index;
   };
 
+  using MapperPath =
+      std::pair<mlir::Operation *, llvm::SmallVector<int64_t, 4>>;
+  using MapperPathStack = llvm::SmallVector<MapperPath, 8>;
+
   /// Tracks any intermediate function/subroutine local allocations we
   /// generate for the descriptors of box type dummy arguments, so that
   /// we can retrieve it for subsequent reuses within the functions
@@ -138,32 +141,47 @@ class MapInfoFinalizationPass
     return findMemberByIndexPath(op, indexPath) != nullptr;
   }
 
-  static bool mapperCoversIndexPath(
-      mlir::Operation *symbolTableAnchor, mlir::FlatSymbolRefAttr mapperId,
-      llvm::ArrayRef<int64_t> indexPath,
-      llvm::SmallPtrSetImpl<mlir::Operation *> &visitedMappers) {
+  static bool mapperCoversIndexPath(mlir::Operation *symbolTableAnchor,
+                                    mlir::FlatSymbolRefAttr mapperId,
+                                    llvm::ArrayRef<int64_t> indexPath,
+                                    MapperPathStack &activeMapperPaths) {
     mlir::omp::DeclareMapperOp symbol =
         mlir::SymbolTable::lookupNearestSymbolFrom<mlir::omp::DeclareMapperOp>(
             symbolTableAnchor, mapperId);
-    if (!symbol || !visitedMappers.insert(symbol.getOperation()).second)
+    if (!symbol)
+      return false;
+
+    mlir::Operation *symbolOp = symbol.getOperation();
+    if (llvm::any_of(activeMapperPaths, [&](const MapperPath &entry) {
+          return entry.first == symbolOp &&
+                 entry.second.size() == indexPath.size() &&
+                 std::equal(entry.second.begin(), entry.second.end(),
+                            indexPath.begin());
+        }))
       return false;
+    activeMapperPaths.emplace_back(
+        symbolOp,
+        llvm::SmallVector<int64_t, 4>(indexPath.begin(), indexPath.end()));
 
     mlir::omp::DeclareMapperInfoOp mapperInfo = symbol.getDeclareMapperInfo();
-    if (!mapperInfo)
+    if (!mapperInfo) {
+      activeMapperPaths.pop_back();
       return false;
+    }
 
-    return llvm::any_of(mapperInfo.getMapVars(), [&](mlir::Value v) {
+    bool covers = llvm::any_of(mapperInfo.getMapVars(), [&](mlir::Value v) {
       mlir::omp::MapInfoOp map =
           mlir::dyn_cast_if_present<mlir::omp::MapInfoOp>(v.getDefiningOp());
-      return map && !map.getMembers().empty() &&
-             map.getMembersIndexAttr() &&
-             mapInfoCoversIndexPath(map, indexPath, visitedMappers);
+      return map && !map.getMembers().empty() && map.getMembersIndexAttr() &&
+             mapInfoCoversIndexPath(map, indexPath, activeMapperPaths);
     });
+    activeMapperPaths.pop_back();
+    return covers;
   }
 
-  static bool mapInfoCoversIndexPath(
-      mlir::omp::MapInfoOp map, llvm::ArrayRef<int64_t> indexPath,
-      llvm::SmallPtrSetImpl<mlir::Operation *> &visitedMappers) {
+  static bool mapInfoCoversIndexPath(mlir::omp::MapInfoOp map,
+                                     llvm::ArrayRef<int64_t> indexPath,
+                                     MapperPathStack &activeMapperPaths) {
     if (mappedIndexPathExists(map, indexPath))
       return true;
 
@@ -195,12 +213,12 @@ class MapInfoFinalizationPass
       llvm::ArrayRef<int64_t> nestedIndexPath =
           indexPath.drop_front(memberIndexPath.size());
       if (!memberMap.getMembers().empty() &&
-          mapInfoCoversIndexPath(memberMap, nestedIndexPath, visitedMappers))
+          mapInfoCoversIndexPath(memberMap, nestedIndexPath, activeMapperPaths))
         return true;
 
       if (memberMap.getMapperIdAttr() &&
           mapperCoversIndexPath(memberMap, memberMap.getMapperIdAttr(),
-                                nestedIndexPath, visitedMappers))
+                                nestedIndexPath, activeMapperPaths))
         return true;
     }
 
@@ -274,9 +292,9 @@ class MapInfoFinalizationPass
       return;
 
     if (op.getMapperId()) {
-      llvm::SmallPtrSet<mlir::Operation *, 4> visitedMappers;
+      MapperPathStack activeMapperPaths;
       if (mapperCoversIndexPath(op, op.getMapperIdAttr(), indexPath,
-                                visitedMappers))
+                                activeMapperPaths))
         return;
     }
 
diff --git a/flang/test/Lower/OpenMP/nested-default-mapper.f90 b/flang/test/Lower/OpenMP/nested-default-mapper.f90
index 1dc0a5bc35aad..6e7ddc8bdd756 100644
--- a/flang/test/Lower/OpenMP/nested-default-mapper.f90
+++ b/flang/test/Lower/OpenMP/nested-default-mapper.f90
@@ -27,6 +27,56 @@ program main
   !$omp end target
 end program main
 
+subroutine recursive_mapper
+  implicit none
+
+  type node_t
+    integer, allocatable :: payload(:)
+    type(node_t), allocatable :: next
+  end type node_t
+
+  type wrapper_t
+    type(node_t), allocatable :: head
+  end type wrapper_t
+
+  type(wrapper_t) :: w
+
+  !$omp target map(tofrom: w)
+    w%head%next%next%payload(1) = 7
+  !$omp end target
+end subroutine recursive_mapper
+
+subroutine mutual_mapper
+  implicit none
+
+  type a_t
+    type(b_t), allocatable :: b
+  end type a_t
+
+  type b_t
+    type(a_t), allocatable :: a
+    integer, allocatable :: payload(:)
+  end type b_t
+
+  type(a_t) :: a
+
+  !$omp target map(tofrom: a)
+    a%b%a%b%payload(1) = 11
+  !$omp end target
+end subroutine mutual_mapper
+
+! CHECK-LABEL: omp.declare_mapper @_QQFmutual_mapper_QFmutual_mapperTa_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} mapper(@_QFmutual_mapperTb_t_omp_default_mapper)
+! CHECK-LABEL: omp.declare_mapper @_QFmutual_mapperTb_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} mapper(@_QQFmutual_mapper_QFmutual_mapperTa_t_omp_default_mapper)
+! CHECK-LABEL: omp.declare_mapper @_QQFmutual_mappera_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} mapper(@_QFmutual_mapperTb_t_omp_default_mapper)
+
+! CHECK-LABEL: omp.declare_mapper @_QFrecursive_mapperTnode_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} mapper(@_QFrecursive_mapperTnode_t_omp_default_mapper)
+! CHECK-LABEL: omp.declare_mapper @_QQFrecursive_mapperwrapper_t_omp_default_mapper
+! CHECK: omp.map.info {{.*}} mapper(@_QFrecursive_mapperTnode_t_omp_default_mapper)
+
 ! CHECK-LABEL: omp.declare_mapper @_QQFtyp_t_omp_default_mapper
 ! CHECK: omp.map.info {{.*}} map_clauses(tofrom) capture(ByRef) mapper(@_QQFnested_t_omp_default_mapper) -> {{.*}} {name = "t%nested"}
 ! CHECK-LABEL: omp.declare_mapper @_QQFnested_t_omp_default_mapper
@@ -35,3 +85,13 @@ end program main
 ! CHECK-NOT: implicit_map
 ! CHECK: %[[TYP_MAP:.*]] = omp.map.info {{.*}} mapper(@_QQFtyp_t_omp_default_mapper){{.*}} {name = "typ"}
 ! CHECK-NEXT: omp.target kernel_type(generic) map_entries(%[[TYP_MAP]] -> %{{[^,]*}} : {{.*}}) {
+
+! CHECK-LABEL: func.func @_QPrecursive_mapper
+! CHECK-NOT: implicit_map
+! CHECK: %[[RECURSIVE_MAP:.*]] = omp.map.info {{.*}} mapper(@_QQFrecursive_mapperwrapper_t_omp_default_mapper){{.*}} {name = "w"}
+! CHECK-NEXT: omp.target kernel_type(generic) map_entries(%[[RECURSIVE_MAP]] -> %{{[^,]*}} : {{.*}}) {
+
+! CHECK-LABEL: func.func @_QPmutual_mapper
+! CHECK-NOT: implicit_map
+! CHECK: %[[MUTUAL_MAP:.*]] = omp.map.info {{.*}} mapper(@_QQFmutual_mappera_t_omp_default_mapper){{.*}} {name = "a"}
+! CHECK-NEXT: omp.target kernel_type(generic) map_entries(%[[MUTUAL_MAP]] -> %{{[^,]*}} : {{.*}}) {

>From 6707780b5e35ab76c1734511f71d9ee171d06e17 Mon Sep 17 00:00:00 2001
From: Akash Banerjee <Akash.Banerjee at amd.com>
Date: Wed, 12 Aug 2026 17:54:43 +0100
Subject: [PATCH 3/3] Add comment.

---
 flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp | 2 ++
 1 file changed, 2 insertions(+)

diff --git a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
index 0340e32d0438f..5d1305b32c0eb 100644
--- a/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
+++ b/flang/lib/Optimizer/OpenMP/MapInfoFinalization.cpp
@@ -189,6 +189,8 @@ class MapInfoFinalizationPass
     if (!memberIndices)
       return false;
 
+    // Match a mapped member whose index path is a prefix of the requested
+    // path, then continue the lookup through that member with the suffix.
     for (auto [memberIdx, memberIndexAttr] : llvm::enumerate(memberIndices)) {
       auto memberIndexPath = mlir::cast<mlir::ArrayAttr>(memberIndexAttr);
       if (memberIndexPath.size() >= indexPath.size())



More information about the flang-commits mailing list