[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
Mon Oct 5 22:51:57 PDT 2026


https://github.com/VijayKandiah created https://github.com/llvm/llvm-project/pull/229313

When the `fir.array_coor` base address comes out of a descriptor, `FIRToMemRef` built each dimension's element stride as a chain of `arith.muli` over the inner extents, even though that descriptor already holds the dimension's cumulative byte stride. This chain grows with the rank while reading the stride will not. This patch reads it from the descriptor instead, as the path for descriptor-owned layouts already does.

The division by the element size is marked `exact`, which lets it lower to a single shift in LLVM. That is valid here because this path handles contiguous storage, where the byte stride is the element size times the product of the inner extents. 

Follow-up: Should the division on the general descriptor path be marked `exact` too? The conversion already needs the byte stride to be a whole multiple of the element size, since memref strides count elements, and it bails out on the derived-type component projections that can break that. I left it alone to keep this patch narrow, but can expand it if that reasoning holds.

>From e5fac16d28a12104385dd3dadec2ca7a14f10a5c 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 1542047724e1d..aad0761296e64 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