[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