[flang-commits] [flang] [flang][OpenMP] Recompute descriptor across workdistribute target fission (PR #225607)

via flang-commits flang-commits at lists.llvm.org
Tue Oct 6 01:48:27 PDT 2026


https://github.com/skc7 updated https://github.com/llvm/llvm-project/pull/225607

>From 983b2dc8f033ac9c7c7bb20f67bc19583856804b Mon Sep 17 00:00:00 2001
From: skc7 <Krishna.Sankisa at amd.com>
Date: Wed, 23 Sep 2026 11:48:38 +0530
Subject: [PATCH 1/2] [flang][OpenMP] Recompute descriptor across
 workdistribute target fission

---
 .../Optimizer/OpenMP/LowerWorkdistribute.cpp  |  6 +++
 ...-workdistribute-fission-recompute-box.mlir | 43 +++++++++++++++++++
 2 files changed, 49 insertions(+)
 create mode 100644 flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir

diff --git a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
index 5af2d1ddb5f50..ec5bd47476485 100644
--- a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
+++ b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
@@ -882,6 +882,12 @@ static bool usedOutsideSplit(Value v, Operation *split) {
 
 /// isRecomputableAfterFission checks if an operation can be recomputed
 static bool isRecomputableAfterFission(Operation *op, Operation *splitBefore) {
+  // A descriptor load must be recomputed from the mapped descriptor in each
+  // split target. Caching the box by value captures a host base_addr that the
+  // flat to/from copy cannot re-attach to the device data.
+  if (auto load = dyn_cast<fir::LoadOp>(op))
+    if (isa<fir::BaseBoxType>(load.getType()))
+      return true;
   // If the op has side effects, it cannot be recomputed.
   // We consider fir.declare as having no side effects.
   return isa<fir::DeclareOp>(op) || isMemoryEffectFree(op);
diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir
new file mode 100644
index 0000000000000..c4ef0dd3ae61c
--- /dev/null
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir
@@ -0,0 +1,43 @@
+// RUN: fir-opt --lower-workdistribute %s | FileCheck %s
+
+// A Fortran allocatable descriptor (fir.box) crossing the workdistribute target
+// fission must be recomputed inside the isolated target from its mapped
+// descriptor, not cached by value via __flang_workdistribute_to/from. Caching
+// the box would freeze a host base_addr into the device kernel.
+
+// CHECK-LABEL: func.func @recompute_descriptor(
+// CHECK: omp.target_data
+// The box must not be cached by value across the split.
+// CHECK-NOT: __flang_workdistribute
+// It is recomputed inside the device target and indexed there.
+// CHECK: fir.load %{{.*}} : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+// CHECK: omp.teams
+// CHECK: omp.loop_nest
+// CHECK: fir.box_addr
+
+module attributes {llvm.target_triple = "amdgcn-amd-amdhsa", omp.is_gpu = true, omp.is_target_device = true} {
+func.func @recompute_descriptor(%arg0: !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+  %map = omp.map.info var_ptr(%arg0 : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(tofrom) capture(ByRef) name("x") -> !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+  omp.target kernel_type(generic) map_entries(%map -> %barg : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+    %c0 = arith.constant 0 : index
+    %c1 = arith.constant 1 : index
+    %c9 = arith.constant 9 : index
+    %cst = arith.constant 5.000000e+00 : f32
+    // The descriptor load crosses the split - it must be recomputed, not cached.
+    %box = fir.load %barg : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    omp.teams {
+      omp.workdistribute {
+        fir.do_loop %iv = %c0 to %c9 step %c1 unordered {
+          %addr = fir.box_addr %box : (!fir.box<!fir.heap<!fir.array<?xf32>>>) -> !fir.heap<!fir.array<?xf32>>
+          %coor = fir.coordinate_of %addr, %iv : (!fir.heap<!fir.array<?xf32>>, index) -> !fir.ref<f32>
+          fir.store %cst to %coor : !fir.ref<f32>
+        }
+        omp.terminator
+      }
+      omp.terminator
+    }
+    omp.terminator
+  }
+  return
+}
+}

>From f009a70fe3e5e9a9fe21959b32163b37ff891c87 Mon Sep 17 00:00:00 2001
From: skc7 <Krishna.Sankisa at amd.com>
Date: Tue, 6 Oct 2026 14:17:31 +0530
Subject: [PATCH 2/2] update

---
 .../Optimizer/OpenMP/LowerWorkdistribute.cpp  | 43 +++++++++++--
 ...-workdistribute-fission-recompute-box.mlir | 60 +++++++++++++++++++
 2 files changed, 99 insertions(+), 4 deletions(-)

diff --git a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
index ec5bd47476485..bd03a192cbc61 100644
--- a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
+++ b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
@@ -880,14 +880,49 @@ static bool usedOutsideSplit(Value v, Operation *split) {
   return false;
 }
 
+/// Returns the map block argument of \p targetOp that \p addr is derived from
+/// through fir.declare and fir.convert, or nullptr otherwise.
+static BlockArgument getMappedBlockArg(Value addr, omp::TargetOp targetOp) {
+  while (Operation *def = addr.getDefiningOp()) {
+    if (auto declare = dyn_cast<fir::DeclareOp>(def))
+      addr = declare.getMemref();
+    else if (auto convert = dyn_cast<fir::ConvertOp>(def))
+      addr = convert.getValue();
+    else
+      return nullptr;
+  }
+  auto argIface = cast<omp::BlockArgOpenMPOpInterface>(*targetOp);
+  auto arg = cast<BlockArgument>(addr);
+  return llvm::is_contained(argIface.getMapBlockArgs(), arg) ? arg : nullptr;
+}
+
+/// Returns true if \p addr, or a fir.declare or fir.convert of it, has a user
+/// other than fir.load.
+static bool hasNonLoadUse(Value addr) {
+  for (Operation *user : addr.getUsers()) {
+    if (isa<fir::LoadOp>(user))
+      continue;
+    if (isa<fir::DeclareOp, fir::ConvertOp>(user) &&
+        !hasNonLoadUse(user->getResult(0)))
+      continue;
+    return true;
+  }
+  return false;
+}
+
 /// isRecomputableAfterFission checks if an operation can be recomputed
 static bool isRecomputableAfterFission(Operation *op, Operation *splitBefore) {
   // A descriptor load must be recomputed from the mapped descriptor in each
   // split target. Caching the box by value captures a host base_addr that the
-  // flat to/from copy cannot re-attach to the device data.
-  if (auto load = dyn_cast<fir::LoadOp>(op))
-    if (isa<fir::BaseBoxType>(load.getType()))
-      return true;
+  // flat to/from copy cannot re-attach to the device data. Only safe when the
+  // mapped descriptor is never written inside the target region.
+  if (auto load = dyn_cast<fir::LoadOp>(op)) {
+    if (isa<fir::BaseBoxType>(load.getType())) {
+      auto targetOp = cast<omp::TargetOp>(splitBefore->getParentOp());
+      BlockArgument arg = getMappedBlockArg(load.getMemref(), targetOp);
+      return arg && !hasNonLoadUse(arg);
+    }
+  }
   // If the op has side effects, it cannot be recomputed.
   // We consider fir.declare as having no side effects.
   return isa<fir::DeclareOp>(op) || isMemoryEffectFree(op);
diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir
index c4ef0dd3ae61c..a3452996e5aa8 100644
--- a/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-fission-recompute-box.mlir
@@ -40,4 +40,64 @@ func.func @recompute_descriptor(%arg0: !fir.ref<!fir.box<!fir.heap<!fir.array<?x
   }
   return
 }
+
+// A box loaded from a local copy must be cached, since the initializing store is not cloned.
+// CHECK-LABEL: func.func @cache_local_descriptor(
+// CHECK: omp.map.info var_ptr({{.*}} : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(from) capture(ByRef) name("__flang_workdistribute_from")
+func.func @cache_local_descriptor(%arg0: !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+  %map = omp.map.info var_ptr(%arg0 : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(tofrom) capture(ByRef) name("x") -> !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+  omp.target kernel_type(generic) map_entries(%map -> %barg : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+    %c0 = arith.constant 0 : index
+    %c1 = arith.constant 1 : index
+    %c9 = arith.constant 9 : index
+    %cst = arith.constant 5.000000e+00 : f32
+    %local = fir.alloca !fir.box<!fir.heap<!fir.array<?xf32>>>
+    %init = fir.load %barg : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    fir.store %init to %local : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    %box = fir.load %local : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    omp.teams {
+      omp.workdistribute {
+        fir.do_loop %iv = %c0 to %c9 step %c1 unordered {
+          %addr = fir.box_addr %box : (!fir.box<!fir.heap<!fir.array<?xf32>>>) -> !fir.heap<!fir.array<?xf32>>
+          %coor = fir.coordinate_of %addr, %iv : (!fir.heap<!fir.array<?xf32>>, index) -> !fir.ref<f32>
+          fir.store %cst to %coor : !fir.ref<f32>
+        }
+        omp.terminator
+      }
+      omp.terminator
+    }
+    omp.terminator
+  }
+  return
+}
+
+// A mapped box that is stored to after the load must be cached to keep the loaded value.
+// CHECK-LABEL: func.func @cache_written_descriptor(
+// CHECK: omp.map.info var_ptr({{.*}} : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(from) capture(ByRef) name("__flang_workdistribute_from")
+func.func @cache_written_descriptor(%arg0: !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, %arg1: !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+  %map0 = omp.map.info var_ptr(%arg0 : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(tofrom) capture(ByRef) name("x") -> !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+  %map1 = omp.map.info var_ptr(%arg1 : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.box<!fir.heap<!fir.array<?xf32>>>) map_clauses(tofrom) capture(ByRef) name("y") -> !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+  omp.target kernel_type(generic) map_entries(%map0 -> %bx, %map1 -> %by : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>, !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) {
+    %c0 = arith.constant 0 : index
+    %c1 = arith.constant 1 : index
+    %c9 = arith.constant 9 : index
+    %cst = arith.constant 5.000000e+00 : f32
+    %box = fir.load %bx : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    %other = fir.load %by : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    fir.store %other to %bx : !fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>
+    omp.teams {
+      omp.workdistribute {
+        fir.do_loop %iv = %c0 to %c9 step %c1 unordered {
+          %addr = fir.box_addr %box : (!fir.box<!fir.heap<!fir.array<?xf32>>>) -> !fir.heap<!fir.array<?xf32>>
+          %coor = fir.coordinate_of %addr, %iv : (!fir.heap<!fir.array<?xf32>>, index) -> !fir.ref<f32>
+          fir.store %cst to %coor : !fir.ref<f32>
+        }
+        omp.terminator
+      }
+      omp.terminator
+    }
+    omp.terminator
+  }
+  return
+}
 }



More information about the flang-commits mailing list