[Mlir-commits] [mlir] [mlir][sparse] Fix mixed sparse-dense iterator lowering (PR #219753)
Kohei Yamaguchi
llvmlistbot at llvm.org
Sat Aug 29 20:39:45 PDT 2026
https://github.com/sott0n updated https://github.com/llvm/llvm-project/pull/219753
>From ceb9393deb6ed46209474a8f17649f0579ef587c Mon Sep 17 00:00:00 2001
From: Kohei Yamaguchi <fix7211 at gmail.com>
Date: Sun, 30 Aug 2026 08:09:34 +0900
Subject: [PATCH] [mlir][sparse] Fix mixed sparse-dense iterator lowering
---
.../Transforms/SparseIterationToScf.cpp | 16 +++--
.../Transforms/Sparsification.cpp | 67 ++++++++++++-------
.../sparse_kernels_to_iterator.mlir | 55 +++++++++++++++
3 files changed, 107 insertions(+), 31 deletions(-)
diff --git a/mlir/lib/Dialect/SparseTensor/Transforms/SparseIterationToScf.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/SparseIterationToScf.cpp
index e4cecef2d44e5..adfb06402f0be 100644
--- a/mlir/lib/Dialect/SparseTensor/Transforms/SparseIterationToScf.cpp
+++ b/mlir/lib/Dialect/SparseTensor/Transforms/SparseIterationToScf.cpp
@@ -228,9 +228,11 @@ class SparseIterateOpConverter : public OpConversionPattern<IterateOp> {
LogicalResult
matchAndRewrite(IterateOp op, OneToNOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
- if (!op.getCrdUsedLvls().empty())
+ I64BitSet crdUsedLvls = op.getCrdUsedLvls();
+ if (crdUsedLvls.count() > 1 ||
+ (!crdUsedLvls.empty() && !crdUsedLvls[op.getSpaceDim() - 1]))
return rewriter.notifyMatchFailure(
- op, "non-empty coordinates list not implemented.");
+ op, "coordinates other than the final level are not implemented");
Location loc = op.getLoc();
@@ -254,13 +256,15 @@ class SparseIterateOpConverter : public OpConversionPattern<IterateOp> {
Block *block = rewriter.applySignatureConversion(
op.getBody(), signatureConversion, getTypeConverter());
+ bool usesCrd = !crdUsedLvls.empty();
ValueRange ret = genLoopWithIterator(
rewriter, loc, it.get(), ivs,
- [block](PatternRewriter &rewriter, Location loc, Region &loopBody,
- SparseIterator *it, ValueRange reduc) -> SmallVector<Value> {
+ [block, usesCrd](PatternRewriter &rewriter, Location loc,
+ Region &loopBody, SparseIterator *it,
+ ValueRange reduc) -> SmallVector<Value> {
SmallVector<Value> blockArgs(reduc);
- // TODO: Also appends coordinates if used.
- // blockArgs.push_back(it->deref(rewriter, loc));
+ if (usesCrd)
+ blockArgs.push_back(it->deref(rewriter, loc));
llvm::append_range(blockArgs, it->getCursor());
Block *dstBlock = &loopBody.getBlocks().front();
diff --git a/mlir/lib/Dialect/SparseTensor/Transforms/Sparsification.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/Sparsification.cpp
index 860ea68029c3a..9b307b31de821 100644
--- a/mlir/lib/Dialect/SparseTensor/Transforms/Sparsification.cpp
+++ b/mlir/lib/Dialect/SparseTensor/Transforms/Sparsification.cpp
@@ -1008,7 +1008,7 @@ static void endIf(CodegenEnv &env, OpBuilder &builder, scf::IfOp ifOp,
static bool getAllTidLvlsInLatPoints(
CodegenEnv &env, LatPointId li, LoopId curr,
- llvm::function_ref<void(TensorLevel, AffineExpr)> callback) {
+ llvm::function_ref<void(TensorLevel, AffineExpr, bool)> callback) {
const BitVector &simple = env.lat(li).simple;
const TensorId outTid = env.merger().getOutTensorID();
const std::optional<Level> outLvl = env.merger().getLvl(outTid, curr);
@@ -1020,7 +1020,8 @@ static bool getAllTidLvlsInLatPoints(
LevelType lt, bool isIdxReduc) {
if (simple[b]) {
if (isIdxReduc) {
- callback(env.makeTensorLevel(tid, *lvl), nullptr);
+ callback(env.makeTensorLevel(tid, *lvl), nullptr,
+ /*isLoopCond=*/true);
numloopCond++;
return;
}
@@ -1044,10 +1045,12 @@ static bool getAllTidLvlsInLatPoints(
}
}
hasNonUnique = !isUniqueLT(lt) || hasNonUnique;
- callback(env.makeTensorLevel(tid, *lvl), nullptr);
+ callback(env.makeTensorLevel(tid, *lvl), nullptr,
+ /*isLoopCond=*/true);
numloopCond++;
} else if (lt.hasDenseSemantic() || isIdxReduc) {
- callback(env.makeTensorLevel(tid, *lvl), nullptr);
+ callback(env.makeTensorLevel(tid, *lvl), nullptr,
+ /*isLoopCond=*/false);
} else {
assert(isUndefLT(lt));
linalg::GenericOp op = env.op();
@@ -1082,7 +1085,8 @@ static bool getAllTidLvlsInLatPoints(
// level. We need to generate the address according to the
// affine expression. This is also the best place we can do it
// to avoid putting it inside inner loops.
- callback(env.makeTensorLevel(tid, l), exp);
+ callback(env.makeTensorLevel(tid, l), exp,
+ /*isLoopCond=*/false);
}
}
}
@@ -1097,14 +1101,16 @@ static bool getAllTidLvlsInLatPoints(
// TODO: we should avoid introducing corner cases for all-dense sparse
// tensors.
if (stt.hasEncoding() && stt.isAllDense())
- callback(env.makeTensorLevel(outTid, *outLvl), nullptr);
+ callback(env.makeTensorLevel(outTid, *outLvl), nullptr,
+ /*isLoopCond=*/false);
}
if (numloopCond == 0) {
// Corner cases where the loop bound is defined by a *unused* operand, in
// this case, we just generate a dense "fake" loop by iterating over the
// synthetic tensor.
- callback(env.makeTensorLevel(env.merger().getSynTensorID(), curr), nullptr);
+ callback(env.makeTensorLevel(env.merger().getSynTensorID(), curr), nullptr,
+ /*isLoopCond=*/true);
numloopCond++;
}
// If we just need to one loop conditions and the conditions is not imposed on
@@ -1129,14 +1135,16 @@ static bool startLoopSeq(CodegenEnv &env, OpBuilder &builder, ExprId exp,
const LatPointId l0 = env.set(lts)[0];
SmallVector<TensorLevel> tidLvls;
- getAllTidLvlsInLatPoints(env, l0, curr, [&](TensorLevel tl, AffineExpr) {
- // TODO: remove this! The same tensor level might be added for multiple
- // times due to the special handling for all-dense "sparse" output tensor
- // (see L1038).
- if (llvm::is_contained(tidLvls, tl))
- return;
- tidLvls.emplace_back(tl);
- });
+ getAllTidLvlsInLatPoints(env, l0, curr,
+ [&](TensorLevel tl, AffineExpr, bool) {
+ // TODO: remove this! The same tensor level might
+ // be added for multiple times due to the special
+ // handling for all-dense "sparse" output tensor
+ // (see L1038).
+ if (llvm::is_contained(tidLvls, tl))
+ return;
+ tidLvls.emplace_back(tl);
+ });
env.emitter().enterNewLoopSeq(builder, env.op().getLoc(), tidLvls);
@@ -1190,14 +1198,18 @@ static void genInitConstantDenseAddress(CodegenEnv &env,
static bool translateBitsToTidLvlPairs(
CodegenEnv &env, LatPointId li, LoopId curr,
SmallVectorImpl<TensorLevel> &tidLvls,
+ SmallVectorImpl<TensorLevel> &condTidLvls,
SmallVectorImpl<std::pair<TensorLevel, AffineExpr>> &affineTidLvls) {
- return getAllTidLvlsInLatPoints(env, li, curr,
- [&](TensorLevel tl, AffineExpr exp) {
- if (exp)
- affineTidLvls.emplace_back(tl, exp);
- else
- tidLvls.emplace_back(tl);
- });
+ return getAllTidLvlsInLatPoints(
+ env, li, curr, [&](TensorLevel tl, AffineExpr exp, bool isLoopCond) {
+ if (exp)
+ affineTidLvls.emplace_back(tl, exp);
+ else {
+ tidLvls.emplace_back(tl);
+ if (isLoopCond)
+ condTidLvls.emplace_back(tl);
+ }
+ });
}
/// Starts a single loop in current sequence.
@@ -1209,15 +1221,20 @@ static std::pair<Operation *, bool> startLoop(CodegenEnv &env,
// after fully migration.
// The set of tensors + lvls to generate loops on
SmallVector<TensorLevel> tidLvls;
+ SmallVector<TensorLevel> condTidLvls;
// The set of dense tensors with non-trivial affine expression that just
// becomes invariant and the address are generated at the current level.
SmallVector<std::pair<TensorLevel, AffineExpr>> affineTidLvls;
- bool isSingleCond =
- translateBitsToTidLvlPairs(env, li, curr, tidLvls, affineTidLvls);
+ bool isSingleCond = translateBitsToTidLvlPairs(env, li, curr, tidLvls,
+ condTidLvls, affineTidLvls);
// Emit the for/while-loop control.
- Operation *loop = genLoop(env, builder, curr, numCases, needsUniv, tidLvls);
+ ArrayRef<TensorLevel> loopTidLvls = env.generatingSparseIterator()
+ ? ArrayRef(condTidLvls)
+ : ArrayRef(tidLvls);
+ Operation *loop =
+ genLoop(env, builder, curr, numCases, needsUniv, loopTidLvls);
Location loc = env.op().getLoc();
for (auto [tidLvl, exp] : affineTidLvls) {
env.emitter().locateLvlAtAffineAddress(builder, loc, tidLvl, exp);
diff --git a/mlir/test/Dialect/SparseTensor/sparse_kernels_to_iterator.mlir b/mlir/test/Dialect/SparseTensor/sparse_kernels_to_iterator.mlir
index 4ca1351449890..0ad5ca575fd5c 100644
--- a/mlir/test/Dialect/SparseTensor/sparse_kernels_to_iterator.mlir
+++ b/mlir/test/Dialect/SparseTensor/sparse_kernels_to_iterator.mlir
@@ -16,6 +16,10 @@
map = (d0) -> (d0 : compressed)
}>
+#CSR = #sparse_tensor.encoding<{
+ map = (d0, d1) -> (d0 : dense, d1 : compressed)
+}>
+
// CHECK-LABEL: func.func @sqsum(
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
@@ -167,3 +171,54 @@ func.func @add(%arg0: tensor<10xi32, #VEC>, %arg1: tensor<10xi32, #VEC>) -> tens
} -> tensor<10xi32>
return %0 : tensor<10xi32>
}
+
+// ITER-LABEL: func.func @matvec(
+// ITER: %[[VECTOR:.*]] = bufferization.to_buffer
+// ITER: %[[ROW_SPACE:.*]] = sparse_tensor.extract_iteration_space %{{.*}} lvls = 0
+// ITER: sparse_tensor.iterate %[[ROW_ITER:.*]] in %[[ROW_SPACE]] at(%[[ROW:[a-zA-Z0-9_]+]])
+// ITER: %[[COL_SPACE:.*]] = sparse_tensor.extract_iteration_space %{{.*}} at %[[ROW_ITER]] lvls = 1
+// ITER: sparse_tensor.iterate %[[COL_ITER:.*]] in %[[COL_SPACE]] at(%[[COL:[a-zA-Z0-9_]+]])
+// ITER: %[[MATRIX_VALUE:.*]] = sparse_tensor.extract_value %{{.*}} at %[[COL_ITER]]
+// ITER: %[[VECTOR_VALUE:.*]] = memref.load %[[VECTOR]][%[[COL]]]
+// ITER: %[[PRODUCT:.*]] = arith.mulf %[[MATRIX_VALUE]], %[[VECTOR_VALUE]] : f64
+// ITER: arith.addf %{{.*}}, %[[PRODUCT]] : f64
+
+// CHECK-LABEL: func.func @matvec(
+// CHECK: %[[C0:.*]] = arith.constant 0 : index
+// CHECK: %[[C1:.*]] = arith.constant 1 : index
+// CHECK: %[[VECTOR:.*]] = bufferization.to_buffer
+// CHECK: %[[OUTPUT:.*]] = bufferization.to_buffer
+// CHECK: scf.for %[[ROW:.*]] = %[[C0]] to %{{.*}} step %[[C1]] {
+// CHECK: %[[SUM:.*]] = memref.load %[[OUTPUT]][%[[ROW]]]
+// CHECK: %[[COL_LO:.*]] = memref.load %{{.*}}[%[[ROW]]]
+// CHECK: %[[ROW_END:.*]] = arith.addi %[[ROW]], %[[C1]] : index
+// CHECK: %[[COL_HI:.*]] = memref.load %{{.*}}[%[[ROW_END]]]
+// CHECK: %[[RESULT:.*]] = scf.for %[[POS:.*]] = %[[COL_LO]] to %[[COL_HI]] step %[[C1]] iter_args(%[[ACC:.*]] = %[[SUM]]) -> (f64) {
+// CHECK: %[[COL:.*]] = memref.load %{{.*}}[%[[POS]]]
+// CHECK: %[[MATRIX_VALUE:.*]] = memref.load %{{.*}}[%[[POS]]]
+// CHECK: %[[VECTOR_VALUE:.*]] = memref.load %[[VECTOR]][%[[COL]]]
+// CHECK: %[[PRODUCT:.*]] = arith.mulf %[[MATRIX_VALUE]], %[[VECTOR_VALUE]] : f64
+// CHECK: %[[NEXT:.*]] = arith.addf %[[ACC]], %[[PRODUCT]] : f64
+// CHECK: scf.yield %[[NEXT]] : f64
+// CHECK: }
+// CHECK: memref.store %[[RESULT]], %[[OUTPUT]][%[[ROW]]]
+// CHECK: }
+func.func @matvec(%matrix: tensor<4x8xf64, #CSR>,
+ %vector: tensor<8xf64>,
+ %output: tensor<4xf64>) -> tensor<4xf64> {
+ %result = linalg.generic {
+ indexing_maps = [
+ affine_map<(i, j) -> (i, j)>,
+ affine_map<(i, j) -> (j)>,
+ affine_map<(i, j) -> (i)>
+ ],
+ iterator_types = ["parallel", "reduction"]
+ } ins(%matrix, %vector : tensor<4x8xf64, #CSR>, tensor<8xf64>)
+ outs(%output : tensor<4xf64>) {
+ ^bb0(%matrixValue: f64, %vectorValue: f64, %sum: f64):
+ %product = arith.mulf %matrixValue, %vectorValue : f64
+ %next = arith.addf %sum, %product : f64
+ linalg.yield %next : f64
+ } -> tensor<4xf64>
+ return %result : tensor<4xf64>
+}
More information about the Mlir-commits
mailing list