[flang-commits] [flang] [flang][FIRToMemRef] Derive element strides from descriptor byte strides (PR #229313)
Vijay Kandiah via flang-commits
flang-commits at lists.llvm.org
Wed Oct 7 10:47:14 PDT 2026
https://github.com/VijayKandiah updated https://github.com/llvm/llvm-project/pull/229313
>From 6fafab6c4760114e6df731b9d5f0b3ca572dbb9e Mon Sep 17 00:00:00 2001
From: Vijay Kandiah <vkandiah at nvidia.com>
Date: Mon, 5 Oct 2026 22:49:08 -0700
Subject: [PATCH] [flang][FIRToMemRef] Derive element strides from descriptor
byte strides
---
.../lib/Optimizer/Transforms/FIRToMemRef.cpp | 40 ++++++++++++++++--
.../array-coor-box-addr-strides.mlir | 42 +++++++++++++++++++
2 files changed, 79 insertions(+), 3 deletions(-)
create mode 100644 flang/test/Transforms/FIRToMemRef/array-coor-box-addr-strides.mlir
diff --git a/flang/lib/Optimizer/Transforms/FIRToMemRef.cpp b/flang/lib/Optimizer/Transforms/FIRToMemRef.cpp
index fd22c2b789aaa..5be6ec35a3972 100644
--- a/flang/lib/Optimizer/Transforms/FIRToMemRef.cpp
+++ b/flang/lib/Optimizer/Transforms/FIRToMemRef.cpp
@@ -916,6 +916,17 @@ FIRToMemRef::getMemrefIndices(fir::ArrayCoorOp arrayCoorOp, Operation *memref,
return indices;
}
+/// Converts a descriptor's byte stride into the element stride a memref needs.
+/// `isExact` marks the division with `exact`.
+static Value elementStrideFromByteStride(Value byteStride, Value elementSize,
+ PatternRewriter &rewriter,
+ Location loc, bool isExact) {
+ auto div = arith::DivSIOp::create(rewriter, loc, byteStride, elementSize);
+ if (isExact)
+ div.setIsExact(true);
+ return castTypeToIndexType(div, rewriter);
+}
+
MemRefInfo
FIRToMemRef::convertArrayCoorOp(Operation *memOp, fir::ArrayCoorOp arrayCoorOp,
PatternRewriter &rewriter,
@@ -1105,9 +1116,8 @@ FIRToMemRef::convertArrayCoorOp(Operation *memOp, fir::ArrayCoorOp arrayCoorOp,
sizes.push_back(castTypeToIndexType(extent, rewriter));
Value byteStride = boxDims->getResult(2);
- Value div =
- arith::DivSIOp::create(rewriter, loc, byteStride, boxElementSize);
- strides.push_back(castTypeToIndexType(div, rewriter));
+ strides.push_back(elementStrideFromByteStride(
+ byteStride, boxElementSize, rewriter, loc, /*isExact=*/false));
}
} else {
@@ -1152,11 +1162,35 @@ FIRToMemRef::convertArrayCoorOp(Operation *memOp, fir::ArrayCoorOp arrayCoorOp,
const bool hasParentShape = firMemrefIsEmbox && arrayCoorOp.getSlice() &&
shapeVec.size() >= acRank + rank;
const unsigned parentShapeStartIdx = hasParentShape ? acRank : 0;
+
+ // A descriptor behind the base address already holds each dimension's
+ // cumulative stride. Read it from the descriptor.
+ Value layoutDescriptor;
+ if (!hasParentShape && !arrayCoorOp.getSlice() && !complexPartIdx) {
+ auto boxAddr = firMemref.getDefiningOp<fir::BoxAddrOp>();
+ if (boxAddr && mlir::isa<fir::BaseBoxType>(boxAddr.getVal().getType()))
+ layoutDescriptor = boxAddr.getVal();
+ }
+ Value layoutEleSize;
+ if (layoutDescriptor)
+ layoutEleSize =
+ fir::BoxEleSizeOp::create(rewriter, loc, indexTy, layoutDescriptor);
+
for (unsigned i = rank - 1; i > 0; --i) {
// Sizes are always the box/slice's visible extents (shapeVec[0..rank-1]).
Value size = shapeVec[i];
sizes.push_back(castTypeToIndexType(size, rewriter));
+ if (layoutDescriptor) {
+ Value dim = arith::ConstantIndexOp::create(rewriter, loc, i);
+ auto boxDims = fir::BoxDimsOp::create(rewriter, loc, indexTy, indexTy,
+ indexTy, layoutDescriptor, dim);
+ strides.push_back(elementStrideFromByteStride(
+ boxDims->getResult(2), layoutEleSize, rewriter, loc,
+ /*isExact=*/true));
+ continue;
+ }
+
// Strides use the parent's extents (via `parentShapeStartIdx`).
Value stride = shapeVec[parentShapeStartIdx + 0];
for (unsigned j = 1; j <= i - 1; ++j)
diff --git a/flang/test/Transforms/FIRToMemRef/array-coor-box-addr-strides.mlir b/flang/test/Transforms/FIRToMemRef/array-coor-box-addr-strides.mlir
new file mode 100644
index 0000000000000..1c24751d2d493
--- /dev/null
+++ b/flang/test/Transforms/FIRToMemRef/array-coor-box-addr-strides.mlir
@@ -0,0 +1,42 @@
+// RUN: fir-opt %s --fir-to-memref --allow-unregistered-dialect | FileCheck %s
+
+// The base address comes from a descriptor, so each dimension's element stride
+// is its byte stride divided by the element size, not a product of the inner
+// extents. The storage is contiguous, so the division is exact.
+// CHECK-LABEL: func.func @box_addr_strides_from_descriptor
+// CHECK: [[ESIZE:%[0-9]+]] = fir.box_elesize
+// CHECK: [[DIMS2:%[0-9]+]]:3 = fir.box_dims
+// CHECK: [[STR2:%[0-9]+]] = arith.divsi [[DIMS2]]#2, [[ESIZE]] exact
+// CHECK: [[DIMS1:%[0-9]+]]:3 = fir.box_dims
+// CHECK: [[STR1:%[0-9]+]] = arith.divsi [[DIMS1]]#2, [[ESIZE]] exact
+// CHECK: memref.reinterpret_cast {{.+}}strides: {{\[}}[[STR2]], [[STR1]], %c1{{[_0-9]*}}]
+// CHECK-NOT: fir.array_coor
+func.func @box_addr_strides_from_descriptor(%box: !fir.box<!fir.heap<!fir.array<?x?x?xf32>>>, %i: index, %j: index, %k: index) {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c2 = arith.constant 2 : index
+ %cst = arith.constant 1.0 : f32
+ %addr = fir.box_addr %box : (!fir.box<!fir.heap<!fir.array<?x?x?xf32>>>) -> !fir.heap<!fir.array<?x?x?xf32>>
+ %d0:3 = fir.box_dims %box, %c0 : (!fir.box<!fir.heap<!fir.array<?x?x?xf32>>>, index) -> (index, index, index)
+ %d1:3 = fir.box_dims %box, %c1 : (!fir.box<!fir.heap<!fir.array<?x?x?xf32>>>, index) -> (index, index, index)
+ %d2:3 = fir.box_dims %box, %c2 : (!fir.box<!fir.heap<!fir.array<?x?x?xf32>>>, index) -> (index, index, index)
+ %shape = fir.shape_shift %d0#0, %d0#1, %d1#0, %d1#1, %d2#0, %d2#1 : (index, index, index, index, index, index) -> !fir.shapeshift<3>
+ %elem = fir.array_coor %addr(%shape) %i, %j, %k : (!fir.heap<!fir.array<?x?x?xf32>>, !fir.shapeshift<3>, index, index, index) -> !fir.ref<f32>
+ fir.store %cst to %elem : !fir.ref<f32>
+ return
+}
+
+// No descriptor behind the base, so the outer stride is still the product of
+// the inner extents.
+// CHECK-LABEL: func.func @raw_ref_strides_from_extents
+// CHECK: [[MUL:%[0-9]+]] = arith.muli %arg{{[0-9]+}}, %arg{{[0-9]+}} : index
+// CHECK: memref.reinterpret_cast {{.+}}strides: {{\[}}[[MUL]], %arg{{[0-9]+}}, %c1{{[_0-9]*}}]
+// CHECK-NOT: fir.box_dims
+// CHECK-NOT: fir.array_coor
+func.func @raw_ref_strides_from_extents(%ref: !fir.ref<!fir.array<?x?x?xf32>>, %e0: index, %e1: index, %e2: index, %i: index, %j: index, %k: index) {
+ %cst = arith.constant 1.0 : f32
+ %shape = fir.shape %e0, %e1, %e2 : (index, index, index) -> !fir.shape<3>
+ %elem = fir.array_coor %ref(%shape) %i, %j, %k : (!fir.ref<!fir.array<?x?x?xf32>>, !fir.shape<3>, index, index, index) -> !fir.ref<f32>
+ fir.store %cst to %elem : !fir.ref<f32>
+ return
+}
More information about the flang-commits
mailing list