[flang-commits] [flang] [flang][OpenMP] Lower array-section workdistribute assign to element loop (PR #225595)

via flang-commits flang-commits at lists.llvm.org
Mon Oct 5 22:56:27 PDT 2026


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

>From 07c6b1f6b9fee9cb0fb9b2a2b281653ddaee3666 Mon Sep 17 00:00:00 2001
From: skc7 <Krishna.Sankisa at amd.com>
Date: Wed, 23 Sep 2026 09:59:44 +0530
Subject: [PATCH 1/3] [flang][OpenMP] Lower array-section workdistribute assign
 to element loop

---
 .../Optimizer/OpenMP/LowerWorkdistribute.cpp  | 122 ++++++++++++++++--
 ...r-workdistribute-runtime-assign-array.mlir |  69 ++++++++++
 2 files changed, 180 insertions(+), 11 deletions(-)
 create mode 100644 flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir

diff --git a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
index 5af2d1ddb5f50..0be7125ea7bf7 100644
--- a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
+++ b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
@@ -445,6 +445,13 @@ static bool isEnclosedTypeBoxScalar(Type type) {
   return false;
 }
 
+/// Check if the type is a fir.box that encloses an array (by value, no ref).
+static bool isEnclosedTypeBoxArray(Type type) {
+  if (auto boxType = dyn_cast<fir::BoxType>(type))
+    return isa<fir::SequenceType>(boxType.getEleTy());
+  return false;
+}
+
 /// Check if the FortranAAssign call has src as scalar and dest as array
 static bool isFortranAssignSrcScalarAndDestArray(fir::CallOp callOp) {
   if (callOp.getNumOperands() < 2)
@@ -469,6 +476,21 @@ static bool isFortranAssignSrcScalarAndDestArray(fir::CallOp callOp) {
   return srcIsScalar && destIsArray;
 }
 
+/// Check if the FortranAAssign call has both src and dest as array descriptors.
+/// This is the array-section copy case (e.g. a(0,:,:) = b(n,:,:)) that a flat
+/// omp_target_memcpy would get wrong for strided sections. Matching the runtime
+/// signature, dest is a reference to a box and src is a box passed by value.
+static bool isFortranAssignSrcArrayAndDestArray(fir::CallOp callOp) {
+  if (callOp.getNumOperands() < 2)
+    return false;
+  auto srcConvert = callOp.getOperand(1).getDefiningOp<fir::ConvertOp>();
+  auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
+  if (!srcConvert || !destConvert)
+    return false;
+  return isEnclosedTypeBoxArray(srcConvert.getValue().getType()) &&
+         isEnclosedTypeRefToBoxArray(destConvert.getValue().getType());
+}
+
 /// Convert a flat index to multi-dimensional indices for an array box
 /// Example: 2D array with shape (2,4)
 ///         Col 1  Col 2  Col 3  Col 4
@@ -536,11 +558,11 @@ static Value CalculateTotalElements(OpBuilder &builder, Location loc,
   return totalElems;
 }
 
-/// Replace the FortranAAssign runtime call with an unordered do loop
-static void replaceWithUnorderedDoLoop(OpBuilder &builder, Location loc,
-                                       omp::TeamsOp teamsOp,
-                                       omp::WorkdistributeOp workdistribute,
-                                       fir::CallOp callOp) {
+/// Replace a scalar-to-array FortranAAssign (broadcast) runtime call with an
+/// unordered do loop that stores the scalar into every element.
+static void replaceScalarToArrayAssignWithUnorderedDoLoop(
+    OpBuilder &builder, Location loc, omp::TeamsOp teamsOp,
+    omp::WorkdistributeOp workdistribute, fir::CallOp callOp) {
   auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
   auto srcConvert = callOp.getOperand(1).getDefiningOp<fir::ConvertOp>();
 
@@ -595,10 +617,70 @@ static void replaceWithUnorderedDoLoop(OpBuilder &builder, Location loc,
   fir::StoreOp::create(builder, loc, scalar, elemPtr);
 }
 
+/// Return the array descriptor (fir.box) value behind a FortranAAssign arg.
+/// The arg is the address of a descriptor temp: prefer the box value stored
+/// into it, otherwise load the reference.
+static Value getAssignArrayBox(OpBuilder &builder, Location loc, Value box) {
+  if (auto alloca = box.getDefiningOp<fir::AllocaOp>()) {
+    for (auto *user : alloca->getUsers())
+      if (auto storeOp = dyn_cast<fir::StoreOp>(user)) {
+        box = storeOp.getValue();
+        break;
+      }
+  }
+  if (isa<fir::ReferenceType>(box.getType()))
+    box = fir::LoadOp::create(builder, loc, box);
+  return box;
+}
+
+/// Replace an array-to-array FortranAAssign runtime call with an unordered do
+/// loop that copies element by element. Addressing goes through fir.array_coor
+/// on both descriptors, so each section's own strides and bounds are honored -
+/// unlike a flat memcpy, this is correct for strided sections.
+static void replaceArrayToArrayAssignWithUnorderedDoLoop(
+    OpBuilder &builder, Location loc, omp::TeamsOp teamsOp,
+    omp::WorkdistributeOp workdistribute, fir::CallOp callOp) {
+  auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
+  auto srcConvert = callOp.getOperand(1).getDefiningOp<fir::ConvertOp>();
+
+  builder.setInsertionPoint(teamsOp);
+  Value destBox = getAssignArrayBox(builder, loc, destConvert.getValue());
+  Value srcBox = getAssignArrayBox(builder, loc, srcConvert.getValue());
+
+  // Element type comes from the destination sequence.
+  auto destBoxType = cast<fir::BoxType>(destBox.getType());
+  auto destSeqType = cast<fir::SequenceType>(destBoxType.getEleTy());
+  Type eleTy = destSeqType.getEleTy();
+  auto eleRefTy = fir::ReferenceType::get(eleTy);
+
+  auto c0 = arith::ConstantIndexOp::create(builder, loc, 0);
+  auto c1 = arith::ConstantIndexOp::create(builder, loc, 1);
+  Value totalElems = CalculateTotalElements(builder, loc, destBox);
+
+  auto *workdistributeBlock = &workdistribute.getRegion().front();
+  builder.setInsertionPointToStart(workdistributeBlock);
+  // Single flattened loop: dest and src conform, so one index set fits both.
+  auto doLoop = fir::DoLoopOp::create(builder, loc, c0, totalElems, c1, true);
+  builder.setInsertionPointToStart(doLoop.getBody());
+
+  auto flatIdx = doLoop.getRegion().front().getArgument(0);
+  SmallVector<Value> indices =
+      convertFlatToMultiDim(builder, loc, flatIdx, destBox);
+
+  auto srcPtr =
+      fir::ArrayCoorOp::create(builder, loc, eleRefTy, srcBox, nullptr, nullptr,
+                               ValueRange{indices}, ValueRange{});
+  Value value = fir::LoadOp::create(builder, loc, srcPtr);
+  auto destPtr =
+      fir::ArrayCoorOp::create(builder, loc, eleRefTy, destBox, nullptr,
+                               nullptr, ValueRange{indices}, ValueRange{});
+  fir::StoreOp::create(builder, loc, value, destPtr);
+}
+
 /// workdistributeRuntimeCallLower method finds the runtime calls
-/// nested in teams {workdistribute{}} and
-/// lowers FortranAAssign to unordered do loop if src is scalar and dest is
-/// array. Other runtime calls are not handled currently.
+/// nested in teams {workdistribute{}} and lowers FortranAAssign to an
+/// unordered do loop for scalar-to-array and array-to-array assigns.
+/// Unsupported assign shapes error out. Other runtime calls are left as is.
 static FailureOr<bool>
 workdistributeRuntimeCallLower(omp::WorkdistributeOp workdistribute,
                                SetVector<omp::TargetOp> &targetOpsToProcess) {
@@ -620,6 +702,9 @@ workdistributeRuntimeCallLower(omp::WorkdistributeOp workdistribute,
   bool changed = false;
   // Get the target op parent of teams
   omp::TargetOp targetOp = dyn_cast<omp::TargetOp>(teams->getParentOp());
+  // Runtime-call lowering only applies inside omp.target.
+  if (!targetOp)
+    return false;
   SmallVector<Operation *> opsToErase;
   for (auto &op : workdistribute.getOps()) {
     if (isRuntimeCall(&op)) {
@@ -627,13 +712,28 @@ workdistributeRuntimeCallLower(omp::WorkdistributeOp workdistribute,
       fir::CallOp runtimeCall = cast<fir::CallOp>(op);
       auto funcName = runtimeCall.getCallee()->getRootReference().getValue();
       if (isFortranAssignCall(funcName)) {
-        if (isFortranAssignSrcScalarAndDestArray(runtimeCall) && targetOp) {
+        if (isFortranAssignSrcScalarAndDestArray(runtimeCall)) {
           // Record the target ops to process later
           targetOpsToProcess.insert(targetOp);
-          replaceWithUnorderedDoLoop(rewriter, loc, teams, workdistribute,
-                                     runtimeCall);
+          replaceScalarToArrayAssignWithUnorderedDoLoop(
+              rewriter, loc, teams, workdistribute, runtimeCall);
+          opsToErase.push_back(&op);
+          changed = true;
+        } else if (isFortranAssignSrcArrayAndDestArray(runtimeCall)) {
+          // Array-section copy: element-wise loop honors strides, unlike the
+          // flat omp_target_memcpy fallback used otherwise.
+          targetOpsToProcess.insert(targetOp);
+          replaceArrayToArrayAssignWithUnorderedDoLoop(
+              rewriter, loc, teams, workdistribute, runtimeCall);
           opsToErase.push_back(&op);
           changed = true;
+        } else {
+          // Recognized runtime call, but its argument shape has no lowering.
+          emitError(runtimeCall->getLoc(),
+                    "Runtime call " + funcName +
+                        " with this argument shape is not supported in "
+                        "workdistribute yet.\n");
+          return failure();
         }
       }
     }
diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
new file mode 100644
index 0000000000000..f3f7741e66949
--- /dev/null
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
@@ -0,0 +1,69 @@
+// RUN: fir-opt --lower-workdistribute %s | FileCheck %s
+
+// An array-to-array _FortranAAssign in target teams workdistribute must lower
+// to an element-wise fir.array_coor copy, not a flat omp_target_memcpy.
+
+// Example Fortran code:
+// !$omp target teams workdistribute
+// a(:,:) = b(:,:)
+// !$omp end target teams workdistribute
+
+// CHECK-LABEL:   func.func @array_assign(
+// CHECK:           omp.target_data
+// CHECK:           omp.target
+// CHECK:             omp.teams
+// CHECK:               omp.parallel
+// CHECK:                 omp.distribute
+// CHECK:                   omp.wsloop
+// CHECK:                     omp.loop_nest
+// CHECK:                       %[[SRC:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
+// CHECK:                       %[[VAL:.*]] = fir.load %[[SRC]] : !fir.ref<f32>
+// CHECK:                       %[[DST:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
+// CHECK:                       fir.store %[[VAL]] to %[[DST]] : !fir.ref<f32>
+// CHECK-NOT:         omp_target_memcpy
+
+module attributes {llvm.target_triple = "amdgcn-amd-amdhsa", omp.is_gpu = true, omp.is_target_device = true} {
+func.func @array_assign(%a : !fir.ref<!fir.array<?x?xf32>>, %b : !fir.ref<!fir.array<?x?xf32>>) {
+  %c0 = arith.constant 0 : index
+  %c1 = arith.constant 1 : index
+  %c10 = arith.constant 10 : index
+  %c20 = arith.constant 20 : index
+  %ub0 = arith.subi %c10, %c1 : index
+  %bnd0 = omp.map.bounds lower_bound(%c0 : index) upper_bound(%ub0 : index) extent(%c10 : index) stride(%c1 : index) start_idx(%c1 : index)
+  %ub1 = arith.subi %c20, %c1 : index
+  %bnd1 = omp.map.bounds lower_bound(%c0 : index) upper_bound(%ub1 : index) extent(%c20 : index) stride(%c1 : index) start_idx(%c1 : index)
+  %mapa = omp.map.info var_ptr(%a : !fir.ref<!fir.array<?x?xf32>>, f32) map_clauses(implicit, tofrom) capture(ByRef) bounds(%bnd0, %bnd1) name("a") -> !fir.ref<!fir.array<?x?xf32>>
+  %mapb = omp.map.info var_ptr(%b : !fir.ref<!fir.array<?x?xf32>>, f32) map_clauses(implicit, tofrom) capture(ByRef) bounds(%bnd0, %bnd1) name("b") -> !fir.ref<!fir.array<?x?xf32>>
+  omp.target kernel_type(generic) map_entries(%mapa -> %arga, %mapb -> %argb : !fir.ref<!fir.array<?x?xf32>>, !fir.ref<!fir.array<?x?xf32>>) {
+    // omp.target is isolated from above, so re-declare the extents here.
+    %e0 = arith.constant 10 : index
+    %e1 = arith.constant 20 : index
+    %shape = fir.shape %e0, %e1 : (index, index) -> !fir.shape<2>
+    %da = fir.declare %arga(%shape) {uniq_name = "a"} : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
+    %db = fir.declare %argb(%shape) {uniq_name = "b"} : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
+    omp.teams {
+      %dtmp = fir.alloca !fir.box<!fir.array<?x?xf32>> {pinned}
+      omp.workdistribute {
+        %srcbox = fir.embox %db(%shape) : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.box<!fir.array<?x?xf32>>
+        %dstbox = fir.embox %da(%shape) : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.box<!fir.array<?x?xf32>>
+        fir.store %dstbox to %dtmp : !fir.ref<!fir.box<!fir.array<?x?xf32>>>
+        %str = fir.address_of(@_QQcl) : !fir.ref<!fir.char<1,2>>
+        %line = arith.constant 13 : i32
+        %destc = fir.convert %dtmp : (!fir.ref<!fir.box<!fir.array<?x?xf32>>>) -> !fir.ref<!fir.box<none>>
+        %srcc = fir.convert %srcbox : (!fir.box<!fir.array<?x?xf32>>) -> !fir.box<none>
+        %strc = fir.convert %str : (!fir.ref<!fir.char<1,2>>) -> !fir.ref<i8>
+        fir.call @_FortranAAssignSimple(%destc, %srcc, %strc, %line) : (!fir.ref<!fir.box<none>>, !fir.box<none>, !fir.ref<i8>, i32) -> ()
+        omp.terminator
+      }
+      omp.terminator
+    }
+    omp.terminator
+  }
+  return
+}
+func.func private @_FortranAAssignSimple(!fir.ref<!fir.box<none>>, !fir.box<none>, !fir.ref<i8>, i32) attributes {fir.runtime}
+fir.global linkonce @_QQcl constant : !fir.char<1,2> {
+  %0 = fir.string_lit "f\00"(2) : !fir.char<1,2>
+  fir.has_value %0 : !fir.char<1,2>
+}
+}

>From 211104029d703a3489e599b18ffaa089ec9d97ed Mon Sep 17 00:00:00 2001
From: skc7 <Krishna.Sankisa at amd.com>
Date: Mon, 5 Oct 2026 20:37:12 +0530
Subject: [PATCH 2/3] update

---
 .../Optimizer/OpenMP/LowerWorkdistribute.cpp  | 130 +++++++++++-------
 ...stribute-runtime-assign-array-overlap.mlir |  72 ++++++++++
 ...r-workdistribute-runtime-assign-array.mlir |  49 ++++---
 3 files changed, 188 insertions(+), 63 deletions(-)
 create mode 100644 flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array-overlap.mlir

diff --git a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
index 0be7125ea7bf7..b07f77e8cd9ea 100644
--- a/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
+++ b/flang/lib/Optimizer/OpenMP/LowerWorkdistribute.cpp
@@ -452,6 +452,18 @@ static bool isEnclosedTypeBoxArray(Type type) {
   return false;
 }
 
+/// Element type of a boxed array, peeling an optional enclosing fir.ref.
+/// Returns null if the type is not a (ref to) box of a sequence.
+static Type getBoxArraySeqEleTy(Type type) {
+  if (auto refType = dyn_cast<fir::ReferenceType>(type))
+    type = refType.getEleTy();
+  auto boxType = dyn_cast<fir::BoxType>(type);
+  if (!boxType)
+    return {};
+  auto seqType = dyn_cast<fir::SequenceType>(boxType.getEleTy());
+  return seqType ? seqType.getEleTy() : Type{};
+}
+
 /// Check if the FortranAAssign call has src as scalar and dest as array
 static bool isFortranAssignSrcScalarAndDestArray(fir::CallOp callOp) {
   if (callOp.getNumOperands() < 2)
@@ -487,8 +499,15 @@ static bool isFortranAssignSrcArrayAndDestArray(fir::CallOp callOp) {
   auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
   if (!srcConvert || !destConvert)
     return false;
-  return isEnclosedTypeBoxArray(srcConvert.getValue().getType()) &&
-         isEnclosedTypeRefToBoxArray(destConvert.getValue().getType());
+  Type srcTy = srcConvert.getValue().getType();
+  Type destTy = destConvert.getValue().getType();
+  if (!isEnclosedTypeBoxArray(srcTy) || !isEnclosedTypeRefToBoxArray(destTy))
+    return false;
+  // The raw element load/store only handles trivial, matching element types.
+  // Derived types need finalization and characters need length parameters.
+  Type srcEleTy = getBoxArraySeqEleTy(srcTy);
+  Type destEleTy = getBoxArraySeqEleTy(destTy);
+  return srcEleTy && fir::isa_trivial(srcEleTy) && srcEleTy == destEleTy;
 }
 
 /// Convert a flat index to multi-dimensional indices for an array box
@@ -560,9 +579,9 @@ static Value CalculateTotalElements(OpBuilder &builder, Location loc,
 
 /// Replace a scalar-to-array FortranAAssign (broadcast) runtime call with an
 /// unordered do loop that stores the scalar into every element.
-static void replaceScalarToArrayAssignWithUnorderedDoLoop(
-    OpBuilder &builder, Location loc, omp::TeamsOp teamsOp,
-    omp::WorkdistributeOp workdistribute, fir::CallOp callOp) {
+static void replaceScalarToArrayAssignWithUnorderedDoLoop(OpBuilder &builder,
+                                                          Location loc,
+                                                          fir::CallOp callOp) {
   auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
   auto srcConvert = callOp.getOperand(1).getDefiningOp<fir::ConvertOp>();
 
@@ -585,7 +604,7 @@ static void replaceScalarToArrayAssignWithUnorderedDoLoop(
     }
   }
 
-  builder.setInsertionPoint(teamsOp);
+  builder.setInsertionPoint(callOp);
   // Load destination array box (if it's a reference)
   Value arrayBox = destBox;
   if (isa<fir::ReferenceType>(destBox.getType()))
@@ -598,11 +617,10 @@ static void replaceScalarToArrayAssignWithUnorderedDoLoop(
   auto c0 = arith::ConstantIndexOp::create(builder, loc, 0);
   auto c1 = arith::ConstantIndexOp::create(builder, loc, 1);
   Value totalElems = CalculateTotalElements(builder, loc, arrayBox);
+  // fir.do_loop upper bound is inclusive, so iterate [0, totalElems - 1].
+  Value ub = arith::SubIOp::create(builder, loc, totalElems, c1);
 
-  auto *workdistributeBlock = &workdistribute.getRegion().front();
-  builder.setInsertionPointToStart(workdistributeBlock);
-  // Create single unordered loop for flattened array
-  auto doLoop = fir::DoLoopOp::create(builder, loc, c0, totalElems, c1, true);
+  auto doLoop = fir::DoLoopOp::create(builder, loc, c0, ub, c1, true);
   Block *loopBlock = &doLoop.getRegion().front();
   builder.setInsertionPointToStart(doLoop.getBody());
 
@@ -633,17 +651,18 @@ static Value getAssignArrayBox(OpBuilder &builder, Location loc, Value box) {
   return box;
 }
 
-/// Replace an array-to-array FortranAAssign runtime call with an unordered do
-/// loop that copies element by element. Addressing goes through fir.array_coor
-/// on both descriptors, so each section's own strides and bounds are honored -
-/// unlike a flat memcpy, this is correct for strided sections.
-static void replaceArrayToArrayAssignWithUnorderedDoLoop(
-    OpBuilder &builder, Location loc, omp::TeamsOp teamsOp,
-    omp::WorkdistributeOp workdistribute, fir::CallOp callOp) {
+/// Replace an array-to-array FortranAAssign runtime call with two unordered do
+/// loops that copy src into a heap temporary and then the temporary into dest.
+/// Addressing goes through fir.array_coor on both descriptors, so each
+/// section's own strides and bounds are honored - unlike a flat memcpy, this is
+/// correct for strided sections.
+static void replaceArrayToArrayAssignWithUnorderedDoLoop(OpBuilder &builder,
+                                                         Location loc,
+                                                         fir::CallOp callOp) {
   auto destConvert = callOp.getOperand(0).getDefiningOp<fir::ConvertOp>();
   auto srcConvert = callOp.getOperand(1).getDefiningOp<fir::ConvertOp>();
 
-  builder.setInsertionPoint(teamsOp);
+  builder.setInsertionPoint(callOp);
   Value destBox = getAssignArrayBox(builder, loc, destConvert.getValue());
   Value srcBox = getAssignArrayBox(builder, loc, srcConvert.getValue());
 
@@ -656,25 +675,40 @@ static void replaceArrayToArrayAssignWithUnorderedDoLoop(
   auto c0 = arith::ConstantIndexOp::create(builder, loc, 0);
   auto c1 = arith::ConstantIndexOp::create(builder, loc, 1);
   Value totalElems = CalculateTotalElements(builder, loc, destBox);
+  // fir.do_loop upper bound is inclusive, so iterate [0, totalElems - 1].
+  Value ub = arith::SubIOp::create(builder, loc, totalElems, c1);
+
+  // Fortran evaluates the whole RHS before storing, and src and dest may
+  // overlap. Fission later puts each loop in its own kernel, so every read
+  // finishes before any write.
+  auto tmpTy =
+      fir::SequenceType::get({fir::SequenceType::getUnknownExtent()}, eleTy);
+  Value tmp = fir::AllocMemOp::create(builder, loc, tmpTy, ValueRange{},
+                                      ValueRange{totalElems});
+
+  // Copies between box[indices] and tmp[flatIdx] over the flattened dest index
+  // space. Dest and src conform, so one index set fits both.
+  auto genCopyLoop = [&](Value box, bool toTmp) {
+    builder.setInsertionPoint(callOp);
+    auto doLoop = fir::DoLoopOp::create(builder, loc, c0, ub, c1, true);
+    builder.setInsertionPointToStart(doLoop.getBody());
+    Value flatIdx = doLoop.getInductionVar();
+    SmallVector<Value> indices =
+        convertFlatToMultiDim(builder, loc, flatIdx, destBox);
+    Value boxPtr = fir::ArrayCoorOp::create(
+        builder, loc, eleRefTy, box, nullptr, nullptr, indices, ValueRange{});
+    Value tmpPtr = fir::CoordinateOp::create(builder, loc, eleRefTy, tmp,
+                                             ValueRange{flatIdx});
+    Value from = toTmp ? boxPtr : tmpPtr;
+    Value to = toTmp ? tmpPtr : boxPtr;
+    Value value = fir::LoadOp::create(builder, loc, from);
+    fir::StoreOp::create(builder, loc, value, to);
+  };
+  genCopyLoop(srcBox, /*toTmp=*/true);
+  genCopyLoop(destBox, /*toTmp=*/false);
 
-  auto *workdistributeBlock = &workdistribute.getRegion().front();
-  builder.setInsertionPointToStart(workdistributeBlock);
-  // Single flattened loop: dest and src conform, so one index set fits both.
-  auto doLoop = fir::DoLoopOp::create(builder, loc, c0, totalElems, c1, true);
-  builder.setInsertionPointToStart(doLoop.getBody());
-
-  auto flatIdx = doLoop.getRegion().front().getArgument(0);
-  SmallVector<Value> indices =
-      convertFlatToMultiDim(builder, loc, flatIdx, destBox);
-
-  auto srcPtr =
-      fir::ArrayCoorOp::create(builder, loc, eleRefTy, srcBox, nullptr, nullptr,
-                               ValueRange{indices}, ValueRange{});
-  Value value = fir::LoadOp::create(builder, loc, srcPtr);
-  auto destPtr =
-      fir::ArrayCoorOp::create(builder, loc, eleRefTy, destBox, nullptr,
-                               nullptr, ValueRange{indices}, ValueRange{});
-  fir::StoreOp::create(builder, loc, value, destPtr);
+  builder.setInsertionPoint(callOp);
+  fir::FreeMemOp::create(builder, loc, tmp);
 }
 
 /// workdistributeRuntimeCallLower method finds the runtime calls
@@ -715,16 +749,16 @@ workdistributeRuntimeCallLower(omp::WorkdistributeOp workdistribute,
         if (isFortranAssignSrcScalarAndDestArray(runtimeCall)) {
           // Record the target ops to process later
           targetOpsToProcess.insert(targetOp);
-          replaceScalarToArrayAssignWithUnorderedDoLoop(
-              rewriter, loc, teams, workdistribute, runtimeCall);
+          replaceScalarToArrayAssignWithUnorderedDoLoop(rewriter, loc,
+                                                        runtimeCall);
           opsToErase.push_back(&op);
           changed = true;
         } else if (isFortranAssignSrcArrayAndDestArray(runtimeCall)) {
           // Array-section copy: element-wise loop honors strides, unlike the
           // flat omp_target_memcpy fallback used otherwise.
           targetOpsToProcess.insert(targetOp);
-          replaceArrayToArrayAssignWithUnorderedDoLoop(
-              rewriter, loc, teams, workdistribute, runtimeCall);
+          replaceArrayToArrayAssignWithUnorderedDoLoop(rewriter, loc,
+                                                       runtimeCall);
           opsToErase.push_back(&op);
           changed = true;
         } else {
@@ -1925,27 +1959,29 @@ class LowerWorkdistributePass
     if (verify.wasInterrupted())
       return signalPassFailure();
 
-    auto fission =
+    // Runs before fission so the loops it emits get split into their own
+    // teams regions.
+    auto rtCallLower =
         moduleOp->walk([&](mlir::omp::WorkdistributeOp workdistribute) {
-          auto res = fissionWorkdistribute(workdistribute);
+          auto res = workdistributeRuntimeCallLower(workdistribute,
+                                                    targetOpsToProcess);
           if (failed(res))
             return WalkResult::interrupt();
           changed |= *res;
           return WalkResult::advance();
         });
-    if (fission.wasInterrupted())
+    if (rtCallLower.wasInterrupted())
       return signalPassFailure();
 
-    auto rtCallLower =
+    auto fission =
         moduleOp->walk([&](mlir::omp::WorkdistributeOp workdistribute) {
-          auto res = workdistributeRuntimeCallLower(workdistribute,
-                                                    targetOpsToProcess);
+          auto res = fissionWorkdistribute(workdistribute);
           if (failed(res))
             return WalkResult::interrupt();
           changed |= *res;
           return WalkResult::advance();
         });
-    if (rtCallLower.wasInterrupted())
+    if (fission.wasInterrupted())
       return signalPassFailure();
 
     moduleOp->walk([&](mlir::omp::WorkdistributeOp workdistribute) {
diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array-overlap.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array-overlap.mlir
new file mode 100644
index 0000000000000..92c5260f44ace
--- /dev/null
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array-overlap.mlir
@@ -0,0 +1,72 @@
+// RUN: fir-opt --lower-workdistribute %s | FileCheck %s
+
+// An array-to-array _FortranAAssign whose sections overlap must copy through a
+// temporary. The copy is split into two kernels so every read of the source
+// finishes before any write to the destination.
+
+// Example Fortran code:
+// !$omp target teams workdistribute
+// a(2:10) = a(1:9)
+// !$omp end target teams workdistribute
+
+// CHECK-LABEL:   func.func @overlap_assign(
+// CHECK:           omp.target_data
+// CHECK:           omp.target_allocmem
+// CHECK:           omp.target kernel_type
+// CHECK:           omp.loop_nest
+// CHECK:           %[[SRC:.*]] = fir.array_coor
+// CHECK:           %[[TMP:.*]] = fir.coordinate_of
+// CHECK:           %[[VAL:.*]] = fir.load %[[SRC]]
+// CHECK:           fir.store %[[VAL]] to %[[TMP]]
+// CHECK:           omp.target kernel_type
+// CHECK:           omp.loop_nest
+// CHECK:           %[[DST:.*]] = fir.array_coor
+// CHECK:           %[[TMP2:.*]] = fir.coordinate_of
+// CHECK:           %[[VAL2:.*]] = fir.load %[[TMP2]]
+// CHECK:           fir.store %[[VAL2]] to %[[DST]]
+// CHECK:           omp.target_freemem
+// CHECK-NOT:       fir.call @_FortranAAssign
+
+module attributes {llvm.target_triple = "amdgcn-amd-amdhsa", omp.is_gpu = true, omp.is_target_device = true} {
+func.func @overlap_assign(%a : !fir.ref<!fir.array<?xf32>>) {
+  %c0 = arith.constant 0 : index
+  %c1 = arith.constant 1 : index
+  %c10 = arith.constant 10 : index
+  %ub0 = arith.subi %c10, %c1 : index
+  %bnd0 = omp.map.bounds lower_bound(%c0 : index) upper_bound(%ub0 : index) extent(%c10 : index) stride(%c1 : index) start_idx(%c1 : index)
+  %mapa = omp.map.info var_ptr(%a : !fir.ref<!fir.array<?xf32>>, f32) map_clauses(implicit, tofrom) capture(ByRef) bounds(%bnd0) name("a") -> !fir.ref<!fir.array<?xf32>>
+  omp.target kernel_type(generic) map_entries(%mapa -> %arga : !fir.ref<!fir.array<?xf32>>) {
+    %e0 = arith.constant 10 : index
+    %shape = fir.shape %e0 : (index) -> !fir.shape<1>
+    %da = fir.declare %arga(%shape) uniq_name("a") : (!fir.ref<!fir.array<?xf32>>, !fir.shape<1>) -> !fir.ref<!fir.array<?xf32>>
+    omp.teams {
+      %dtmp = fir.alloca !fir.box<!fir.array<?xf32>> {pinned}
+      omp.workdistribute {
+        %one = arith.constant 1 : index
+        %two = arith.constant 2 : index
+        %nine = arith.constant 9 : index
+        %srcslice = fir.slice %one, %nine, %one : (index, index, index) -> !fir.slice<1>
+        %dstslice = fir.slice %two, %e0, %one : (index, index, index) -> !fir.slice<1>
+        %srcbox = fir.embox %da(%shape) [%srcslice] : (!fir.ref<!fir.array<?xf32>>, !fir.shape<1>, !fir.slice<1>) -> !fir.box<!fir.array<?xf32>>
+        %dstbox = fir.embox %da(%shape) [%dstslice] : (!fir.ref<!fir.array<?xf32>>, !fir.shape<1>, !fir.slice<1>) -> !fir.box<!fir.array<?xf32>>
+        fir.store %dstbox to %dtmp : !fir.ref<!fir.box<!fir.array<?xf32>>>
+        %str = fir.address_of(@_QQcl) : !fir.ref<!fir.char<1,2>>
+        %line = arith.constant 9 : i32
+        %destc = fir.convert %dtmp : (!fir.ref<!fir.box<!fir.array<?xf32>>>) -> !fir.ref<!fir.box<none>>
+        %srcc = fir.convert %srcbox : (!fir.box<!fir.array<?xf32>>) -> !fir.box<none>
+        %strc = fir.convert %str : (!fir.ref<!fir.char<1,2>>) -> !fir.ref<i8>
+        fir.call @_FortranAAssign(%destc, %srcc, %strc, %line) : (!fir.ref<!fir.box<none>>, !fir.box<none>, !fir.ref<i8>, i32) -> ()
+        omp.terminator
+      }
+      omp.terminator
+    }
+    omp.terminator
+  }
+  return
+}
+func.func private @_FortranAAssign(!fir.ref<!fir.box<none>>, !fir.box<none>, !fir.ref<i8>, i32) attributes {fir.runtime}
+fir.global linkonce @_QQcl constant : !fir.char<1,2> {
+  %0 = fir.string_lit "f\00"(2) : !fir.char<1,2>
+  fir.has_value %0 : !fir.char<1,2>
+}
+}
diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
index f3f7741e66949..2ee2570637832 100644
--- a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
@@ -1,7 +1,9 @@
 // RUN: fir-opt --lower-workdistribute %s | FileCheck %s
 
 // An array-to-array _FortranAAssign in target teams workdistribute must lower
-// to an element-wise fir.array_coor copy, not a flat omp_target_memcpy.
+// to an element-wise fir.array_coor copy, not a flat omp_target_memcpy. The
+// copy must address through the (strided) section descriptors. It goes
+// through a heap temporary in two kernels, src to tmp and then tmp to dest.
 
 // Example Fortran code:
 // !$omp target teams workdistribute
@@ -10,17 +12,27 @@
 
 // CHECK-LABEL:   func.func @array_assign(
 // CHECK:           omp.target_data
-// CHECK:           omp.target
-// CHECK:             omp.teams
-// CHECK:               omp.parallel
-// CHECK:                 omp.distribute
-// CHECK:                   omp.wsloop
-// CHECK:                     omp.loop_nest
-// CHECK:                       %[[SRC:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
-// CHECK:                       %[[VAL:.*]] = fir.load %[[SRC]] : !fir.ref<f32>
-// CHECK:                       %[[DST:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
-// CHECK:                       fir.store %[[VAL]] to %[[DST]] : !fir.ref<f32>
-// CHECK-NOT:         omp_target_memcpy
+// CHECK:           omp.target_allocmem
+
+// First kernel reads the strided src section into the temporary.
+// CHECK:           omp.target kernel_type
+// CHECK:           %[[SLICE:.*]] = fir.slice
+// CHECK:           %[[SRCBOX:.*]] = fir.embox {{.*}}[%[[SLICE]]]
+// CHECK:           omp.loop_nest
+// CHECK:           %[[SRC:.*]] = fir.array_coor %[[SRCBOX]] {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
+// CHECK:           %[[TMP:.*]] = fir.coordinate_of {{.*}} -> !fir.ref<f32>
+// CHECK:           %[[VAL:.*]] = fir.load %[[SRC]] : !fir.ref<f32>
+// CHECK:           fir.store %[[VAL]] to %[[TMP]] : !fir.ref<f32>
+
+// Second kernel writes the temporary into the strided dest section.
+// CHECK:           omp.target kernel_type
+// CHECK:           omp.loop_nest
+// CHECK:           %[[DST:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
+// CHECK:           %[[TMP2:.*]] = fir.coordinate_of {{.*}} -> !fir.ref<f32>
+// CHECK:           %[[VAL2:.*]] = fir.load %[[TMP2]] : !fir.ref<f32>
+// CHECK:           fir.store %[[VAL2]] to %[[DST]] : !fir.ref<f32>
+// CHECK:           omp.target_freemem
+// CHECK-NOT:       omp_target_memcpy
 
 module attributes {llvm.target_triple = "amdgcn-amd-amdhsa", omp.is_gpu = true, omp.is_target_device = true} {
 func.func @array_assign(%a : !fir.ref<!fir.array<?x?xf32>>, %b : !fir.ref<!fir.array<?x?xf32>>) {
@@ -39,13 +51,18 @@ func.func @array_assign(%a : !fir.ref<!fir.array<?x?xf32>>, %b : !fir.ref<!fir.a
     %e0 = arith.constant 10 : index
     %e1 = arith.constant 20 : index
     %shape = fir.shape %e0, %e1 : (index, index) -> !fir.shape<2>
-    %da = fir.declare %arga(%shape) {uniq_name = "a"} : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
-    %db = fir.declare %argb(%shape) {uniq_name = "b"} : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
+    %da = fir.declare %arga(%shape) uniq_name("a") : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
+    %db = fir.declare %argb(%shape) uniq_name("b") : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.ref<!fir.array<?x?xf32>>
     omp.teams {
       %dtmp = fir.alloca !fir.box<!fir.array<?x?xf32>> {pinned}
       omp.workdistribute {
-        %srcbox = fir.embox %db(%shape) : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.box<!fir.array<?x?xf32>>
-        %dstbox = fir.embox %da(%shape) : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>) -> !fir.box<!fir.array<?x?xf32>>
+        // Strided sections (stride 2 in dim 0): the descriptors carry the
+        // stride, so a correct lowering must address through them.
+        %lb = arith.constant 1 : index
+        %st = arith.constant 2 : index
+        %slice = fir.slice %lb, %e0, %st, %lb, %e1, %lb : (index, index, index, index, index, index) -> !fir.slice<2>
+        %srcbox = fir.embox %db(%shape) [%slice] : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>, !fir.slice<2>) -> !fir.box<!fir.array<?x?xf32>>
+        %dstbox = fir.embox %da(%shape) [%slice] : (!fir.ref<!fir.array<?x?xf32>>, !fir.shape<2>, !fir.slice<2>) -> !fir.box<!fir.array<?x?xf32>>
         fir.store %dstbox to %dtmp : !fir.ref<!fir.box<!fir.array<?x?xf32>>>
         %str = fir.address_of(@_QQcl) : !fir.ref<!fir.char<1,2>>
         %line = arith.constant 13 : i32

>From fbfe4eed1ecd1d60447d313fd7728e0b87f1fe13 Mon Sep 17 00:00:00 2001
From: skc7 <Krishna.Sankisa at amd.com>
Date: Tue, 6 Oct 2026 11:25:55 +0530
Subject: [PATCH 3/3] test update

---
 ...r-workdistribute-runtime-assign-array.mlir | 21 ++++++++++++-------
 1 file changed, 13 insertions(+), 8 deletions(-)

diff --git a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
index 2ee2570637832..a2c2d5426289e 100644
--- a/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
+++ b/flang/test/Transforms/OpenMP/lower-workdistribute-runtime-assign-array.mlir
@@ -7,27 +7,33 @@
 
 // Example Fortran code:
 // !$omp target teams workdistribute
-// a(:,:) = b(:,:)
+// a(1:10:2,:) = b(1:10:2,:)
 // !$omp end target teams workdistribute
 
 // CHECK-LABEL:   func.func @array_assign(
 // CHECK:           omp.target_data
 // CHECK:           omp.target_allocmem
 
-// First kernel reads the strided src section into the temporary.
+// First kernel reads the stride-2 section of b into the temporary.
 // CHECK:           omp.target kernel_type
-// CHECK:           %[[SLICE:.*]] = fir.slice
-// CHECK:           %[[SRCBOX:.*]] = fir.embox {{.*}}[%[[SLICE]]]
+// CHECK:           %[[DB:.*]] = fir.declare {{.*}} uniq_name("b")
+// CHECK:           %[[STEP:.*]] = arith.constant 2 : index
+// CHECK:           %[[SLICE:.*]] = fir.slice {{[^,]*}}, {{[^,]*}}, %[[STEP]], {{.*}} -> !fir.slice<2>
+// CHECK:           %[[SRCBOX:.*]] = fir.embox %[[DB]]({{.*}}) [%[[SLICE]]]
 // CHECK:           omp.loop_nest
 // CHECK:           %[[SRC:.*]] = fir.array_coor %[[SRCBOX]] {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
 // CHECK:           %[[TMP:.*]] = fir.coordinate_of {{.*}} -> !fir.ref<f32>
 // CHECK:           %[[VAL:.*]] = fir.load %[[SRC]] : !fir.ref<f32>
 // CHECK:           fir.store %[[VAL]] to %[[TMP]] : !fir.ref<f32>
 
-// Second kernel writes the temporary into the strided dest section.
+// Second kernel writes the temporary into the stride-2 section of a.
 // CHECK:           omp.target kernel_type
+// CHECK:           %[[DA:.*]] = fir.declare {{.*}} uniq_name("a")
+// CHECK:           %[[STEP2:.*]] = arith.constant 2 : index
+// CHECK:           %[[SLICE2:.*]] = fir.slice {{[^,]*}}, {{[^,]*}}, %[[STEP2]], {{.*}} -> !fir.slice<2>
+// CHECK:           %[[DSTBOX:.*]] = fir.embox %[[DA]]({{.*}}) [%[[SLICE2]]]
 // CHECK:           omp.loop_nest
-// CHECK:           %[[DST:.*]] = fir.array_coor {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
+// CHECK:           %[[DST:.*]] = fir.array_coor %[[DSTBOX]] {{.*}} : (!fir.box<!fir.array<?x?xf32>>, index, index) -> !fir.ref<f32>
 // CHECK:           %[[TMP2:.*]] = fir.coordinate_of {{.*}} -> !fir.ref<f32>
 // CHECK:           %[[VAL2:.*]] = fir.load %[[TMP2]] : !fir.ref<f32>
 // CHECK:           fir.store %[[VAL2]] to %[[DST]] : !fir.ref<f32>
@@ -56,8 +62,7 @@ func.func @array_assign(%a : !fir.ref<!fir.array<?x?xf32>>, %b : !fir.ref<!fir.a
     omp.teams {
       %dtmp = fir.alloca !fir.box<!fir.array<?x?xf32>> {pinned}
       omp.workdistribute {
-        // Strided sections (stride 2 in dim 0): the descriptors carry the
-        // stride, so a correct lowering must address through them.
+        // Stride-2 sections in dim 0, so the copy must address through the descriptors.
         %lb = arith.constant 1 : index
         %st = arith.constant 2 : index
         %slice = fir.slice %lb, %e0, %st, %lb, %e1, %lb : (index, index, index, index, index, index) -> !fir.slice<2>



More information about the flang-commits mailing list