[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