[Mlir-commits] [mlir] dcbf875 - [mlir][sparse] Avoid vectorizing non-contiguous COO coordinate loads (#211004)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Thu Jul 23 07:33:16 PDT 2026


Author: Federico Bruzzone
Date: 2026-07-23T15:33:10+01:00
New Revision: dcbf87542b26480c365b6388b5accf05d0507626

URL: https://github.com/llvm/llvm-project/commit/dcbf87542b26480c365b6388b5accf05d0507626
DIFF: https://github.com/llvm/llvm-project/commit/dcbf87542b26480c365b6388b5accf05d0507626.diff

LOG: [mlir][sparse] Avoid vectorizing non-contiguous COO coordinate loads (#211004)

`SparseVectorization` assumes direct loop accesses (`a[lo:hi]`) are
contiguous and vectorizes them with `vector.maskedload/maskedstore`.
This is false for `sparse_tensor.coordinates` of a level inside a
trailing AoS COO region, whose buffer is interleaved with other levels:
a silent miscompile.

The true stride is already known from the tensor's encoding, even though
the memref type is still dynamic at this point. Use it to fall back to a
scalar loop when the stride is provably non-unit.

---------

Signed-off-by: Federico Bruzzone <federico.bruzzone.i at gmail.com>

Added: 
    mlir/test/Dialect/SparseTensor/sparse_vector_coo.mlir

Modified: 
    mlir/lib/Dialect/SparseTensor/Transforms/SparseVectorization.cpp

Removed: 
    


################################################################################
diff  --git a/mlir/lib/Dialect/SparseTensor/Transforms/SparseVectorization.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/SparseVectorization.cpp
index 23436a68535fc..c60ca523d81f3 100644
--- a/mlir/lib/Dialect/SparseTensor/Transforms/SparseVectorization.cpp
+++ b/mlir/lib/Dialect/SparseTensor/Transforms/SparseVectorization.cpp
@@ -54,6 +54,50 @@ static bool isInvariantArg(BlockArgument arg, Block *block) {
   return arg.getOwner() != block;
 }
 
+/// Returns true when `mem`'s most minor dimension has a statically known
+/// non-unit stride.
+///
+/// `genVectorLoad/genVectorStore` assume a contiguous
+/// `vector.maskedload/vector.maskedstore` is safe for consecutive loop
+/// indices, which breaks for a strided view extracting one component out
+/// of an interleaved (AoS) COO coordinate buffer.
+///
+/// Example:
+/// A `compressed(nonunique) + singleton` region stores coordinates as
+/// `[row0, col0, row1, col1, ...]`, so a 2-lane masked load of `col[0:2]`
+/// (offset=1) would read physical offsets {1, 2} = `[col0, row1]` instead
+/// of the intended {1, 3} = `[col0, col1]`: a silent miscompile.
+///
+/// NOTE: A stride that can't be proven non-unit by either means is assumed
+/// safe.
+static bool hasKnownNonUnitStride(Value mem) {
+  // sparse_tensor.coordinates isn't lowered to a concrete strided memref
+  // until sparse-tensor-codegen runs, so at this point its type
+  // still has a dynamic stride even when the true stride is already known
+  // from the tensor's encoding -- hence the special case below instead of
+  // trusting the memref type.
+  if (auto toCoords = mem.getDefiningOp<ToCoordinatesOp>()) {
+    SparseTensorType stt = getSparseTensorType(toCoords.getTensor());
+    Level cooStart = stt.getAoSCOOStart();
+    // A single trailing level (lvlRank - cooStart == 1) is not actually
+    // interleaved with anything else, so it degenerates to a contiguous
+    // buffer.
+    if (toCoords.getLevel() >= cooStart)
+      return stt.getLvlRank() - cooStart != 1;
+    return false;
+  }
+
+  auto memTp = dyn_cast<MemRefType>(mem.getType());
+  if (!memTp)
+    return false;
+  SmallVector<int64_t> strides;
+  int64_t offset;
+  if (failed(memTp.getStridesAndOffset(strides, offset)))
+    return false;
+  return !strides.empty() && !ShapedType::isDynamic(strides.back()) &&
+         strides.back() != 1;
+}
+
 /// Constructs vector type for element type.
 static VectorType vectorType(VL vl, Type etp) {
   return VectorType::get(vl.vectorLength, etp, vl.enableVLAVectorization);
@@ -292,6 +336,8 @@ static bool vectorizeSubscripts(PatternRewriter &rewriter, scf::ForOp forOp,
     if (auto load = cast.getDefiningOp<memref::LoadOp>()) {
       if (!innermost)
         return false;
+      if (hasKnownNonUnitStride(load.getMemRef()))
+        return false;
       if (codegen) {
         SmallVector<Value> idxs2(load.getIndices()); // no need to analyze
         Location loc = forOp.getLoc();
@@ -408,6 +454,8 @@ static bool vectorizeExpr(PatternRewriter &rewriter, scf::ForOp forOp, VL vl,
   // a[lo:hi] = ind[lo:hi], where 'lo' denotes the current index
   // and 'hi = lo + vl - 1'.
   if (auto load = dyn_cast<memref::LoadOp>(def)) {
+    if (hasKnownNonUnitStride(load.getMemRef()))
+      return false;
     auto subs = load.getIndices();
     SmallVector<Value> idxs;
     if (vectorizeSubscripts(rewriter, forOp, vl, subs, codegen, vmask, idxs)) {
@@ -582,6 +630,8 @@ static bool vectorizeStmt(PatternRewriter &rewriter, scf::ForOp forOp, VL vl,
     }
   } else if (auto store = dyn_cast<memref::StoreOp>(last)) {
     // Analyze/vectorize store operation.
+    if (hasKnownNonUnitStride(store.getMemRef()))
+      return false;
     auto subs = store.getIndices();
     SmallVector<Value> idxs;
     Value rhs = store.getValue();

diff  --git a/mlir/test/Dialect/SparseTensor/sparse_vector_coo.mlir b/mlir/test/Dialect/SparseTensor/sparse_vector_coo.mlir
new file mode 100644
index 0000000000000..9b13b15a60e87
--- /dev/null
+++ b/mlir/test/Dialect/SparseTensor/sparse_vector_coo.mlir
@@ -0,0 +1,105 @@
+// RUN: mlir-opt %s --sparse-reinterpret-map -sparsification -cse -sparse-vectorization="vl=2" -cse | FileCheck %s
+
+// NOTE: Assertions have been autogenerated by utils/generate-test-checks.py
+
+#SortedCOO = #sparse_tensor.encoding<{
+  map = (d0, d1) -> (d0 : compressed(nonunique), d1 : singleton)
+}>
+
+#trait_index = {
+  indexing_maps = [
+    affine_map<(i,j) -> (i,j)>,  // A
+    affine_map<(i,j) -> (i,j)>   // X (out)
+  ],
+  iterator_types = ["parallel", "parallel"],
+  doc = "X(i,j) = A(i,j) * j"
+}
+
+// CHECK: #[[$ATTR_0:.+]] = #sparse_tensor.encoding<{ map = (d0, d1) -> (d0 : compressed(nonunique), d1 : singleton) }>
+// CHECK-LABEL:   func.func @sparse_index_2d_coo(
+// CHECK-SAME:      %[[ARG0:.*]]: tensor<8x8xi64, #[[$ATTR_0]]>) -> tensor<8x8xi64> {
+// CHECK:           %[[CONSTANT_0:.*]] = arith.constant true
+// CHECK:           %[[CONSTANT_1:.*]] = arith.constant false
+// CHECK:           %[[CONSTANT_2:.*]] = arith.constant 1 : index
+// CHECK:           %[[CONSTANT_3:.*]] = arith.constant 0 : index
+// CHECK:           %[[CONSTANT_4:.*]] = arith.constant 0 : i64
+// CHECK:           %[[EMPTY_0:.*]] = tensor.empty() : tensor<8x8xi64>
+// CHECK:           %[[VALUES_0:.*]] = sparse_tensor.values %[[ARG0]] : tensor<8x8xi64, #[[$ATTR_0]]> to memref<?xi64>
+// CHECK:           %[[TO_BUFFER_0:.*]] = bufferization.to_buffer %[[EMPTY_0]] : tensor<8x8xi64> to memref<8x8xi64>
+// CHECK:           linalg.fill ins(%[[CONSTANT_4]] : i64) outs(%[[TO_BUFFER_0]] : memref<8x8xi64>)
+// CHECK:           %[[POSITIONS_0:.*]] = sparse_tensor.positions %[[ARG0]] {level = 0 : index} : tensor<8x8xi64, #[[$ATTR_0]]> to memref<?xindex>
+// CHECK:           %[[COORDINATES_0:.*]] = sparse_tensor.coordinates %[[ARG0]] {level = 0 : index} : tensor<8x8xi64, #[[$ATTR_0]]> to memref<?xindex, strided<[?], offset: ?>>
+// CHECK:           %[[COORDINATES_1:.*]] = sparse_tensor.coordinates %[[ARG0]] {level = 1 : index} : tensor<8x8xi64, #[[$ATTR_0]]> to memref<?xindex, strided<[?], offset: ?>>
+// CHECK:           %[[LOAD_0:.*]] = memref.load %[[POSITIONS_0]]{{\[}}%[[CONSTANT_3]]] : memref<?xindex>
+// CHECK:           %[[LOAD_1:.*]] = memref.load %[[POSITIONS_0]]{{\[}}%[[CONSTANT_2]]] : memref<?xindex>
+// CHECK:           %[[WHILE_0:.*]] = scf.while (%[[VAL_0:.*]] = %[[LOAD_0]]) : (index) -> index {
+// CHECK:             %[[CMPI_0:.*]] = arith.cmpi ult, %[[VAL_0]], %[[LOAD_1]] : index
+// CHECK:             %[[IF_0:.*]] = scf.if %[[CMPI_0]] -> (i1) {
+// CHECK:               %[[LOAD_2:.*]] = memref.load %[[COORDINATES_0]]{{\[}}%[[LOAD_0]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:               %[[LOAD_3:.*]] = memref.load %[[COORDINATES_0]]{{\[}}%[[VAL_0]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:               %[[CMPI_1:.*]] = arith.cmpi eq, %[[LOAD_2]], %[[LOAD_3]] : index
+// CHECK:               scf.yield %[[CMPI_1]] : i1
+// CHECK:             } else {
+// CHECK:               scf.yield %[[CONSTANT_1]] : i1
+// CHECK:             }
+// CHECK:             scf.condition(%[[IF_0]]) %[[VAL_0]] : index
+// CHECK:           } do {
+// CHECK:           ^bb0(%[[VAL_1:.*]]: index):
+// CHECK:             %[[ADDI_0:.*]] = arith.addi %[[VAL_1]], %[[CONSTANT_2]] : index
+// CHECK:             scf.yield %[[ADDI_0]] : index
+// CHECK:           }
+// CHECK:           %[[WHILE_1:.*]]:2 = scf.while (%[[VAL_2:.*]] = %[[LOAD_0]], %[[VAL_3:.*]] = %[[WHILE_0]]) : (index, index) -> (index, index) {
+// CHECK:             %[[CMPI_2:.*]] = arith.cmpi ult, %[[VAL_2]], %[[LOAD_1]] : index
+// CHECK:             scf.condition(%[[CMPI_2]]) %[[VAL_2]], %[[VAL_3]] : index, index
+// CHECK:           } do {
+// CHECK:           ^bb0(%[[VAL_4:.*]]: index, %[[VAL_5:.*]]: index):
+// CHECK:             %[[LOAD_4:.*]] = memref.load %[[COORDINATES_0]]{{\[}}%[[VAL_4]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:             scf.if %[[CONSTANT_0]] {
+// CHECK:               scf.for %[[VAL_6:.*]] = %[[VAL_4]] to %[[VAL_5]] step %[[CONSTANT_2]] {
+// CHECK:                 %[[LOAD_5:.*]] = memref.load %[[COORDINATES_1]]{{\[}}%[[VAL_6]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:                 %[[LOAD_6:.*]] = memref.load %[[VALUES_0]]{{\[}}%[[VAL_6]]] : memref<?xi64>
+// CHECK:                 %[[INDEX_CAST_0:.*]] = arith.index_cast %[[LOAD_5]] : index to i64
+// CHECK:                 %[[MULI_0:.*]] = arith.muli %[[LOAD_6]], %[[INDEX_CAST_0]] : i64
+// CHECK:                 memref.store %[[MULI_0]], %[[TO_BUFFER_0]]{{\[}}%[[LOAD_4]], %[[LOAD_5]]] : memref<8x8xi64>
+// CHECK:               } {"Emitted from" = "linalg.generic"}
+// CHECK:             } else {
+// CHECK:             }
+// CHECK:             %[[IF_1:.*]]:2 = scf.if %[[CONSTANT_0]] -> (index, index) {
+// CHECK:               %[[WHILE_2:.*]] = scf.while (%[[VAL_7:.*]] = %[[VAL_5]]) : (index) -> index {
+// CHECK:                 %[[CMPI_3:.*]] = arith.cmpi ult, %[[VAL_7]], %[[LOAD_1]] : index
+// CHECK:                 %[[IF_2:.*]] = scf.if %[[CMPI_3]] -> (i1) {
+// CHECK:                   %[[LOAD_7:.*]] = memref.load %[[COORDINATES_0]]{{\[}}%[[VAL_5]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:                   %[[LOAD_8:.*]] = memref.load %[[COORDINATES_0]]{{\[}}%[[VAL_7]]] : memref<?xindex, strided<[?], offset: ?>>
+// CHECK:                   %[[CMPI_4:.*]] = arith.cmpi eq, %[[LOAD_7]], %[[LOAD_8]] : index
+// CHECK:                   scf.yield %[[CMPI_4]] : i1
+// CHECK:                 } else {
+// CHECK:                   scf.yield %[[CONSTANT_1]] : i1
+// CHECK:                 }
+// CHECK:                 scf.condition(%[[IF_2]]) %[[VAL_7]] : index
+// CHECK:               } do {
+// CHECK:               ^bb0(%[[VAL_8:.*]]: index):
+// CHECK:                 %[[ADDI_1:.*]] = arith.addi %[[VAL_8]], %[[CONSTANT_2]] : index
+// CHECK:                 scf.yield %[[ADDI_1]] : index
+// CHECK:               }
+// CHECK:               scf.yield %[[VAL_5]], %[[WHILE_2]] : index, index
+// CHECK:             } else {
+// CHECK:               scf.yield %[[VAL_4]], %[[VAL_5]] : index, index
+// CHECK:             }
+// CHECK:             scf.yield %[[VAL_9:.*]]#0, %[[VAL_9]]#1 : index, index
+// CHECK:           } attributes {"Emitted from" = "linalg.generic"}
+// CHECK:           %[[TO_TENSOR_0:.*]] = bufferization.to_tensor %[[TO_BUFFER_0]] : memref<8x8xi64> to tensor<8x8xi64>
+// CHECK:           return %[[TO_TENSOR_0]] : tensor<8x8xi64>
+// CHECK:         }
+func.func @sparse_index_2d_coo(%arga: tensor<8x8xi64, #SortedCOO>) -> tensor<8x8xi64> {
+  %init = tensor.empty() : tensor<8x8xi64>
+  %r = linalg.generic #trait_index
+      ins(%arga: tensor<8x8xi64, #SortedCOO>)
+     outs(%init: tensor<8x8xi64>) {
+      ^bb(%a: i64, %x: i64):
+        %j = linalg.index 1 : index
+        %jj = arith.index_cast %j : index to i64
+        %m1 = arith.muli %a, %jj : i64
+        linalg.yield %m1 : i64
+  } -> tensor<8x8xi64>
+  return %r : tensor<8x8xi64>
+}


        


More information about the Mlir-commits mailing list