[Mlir-commits] [mlir] [mlir][sparse] Collapse dense and compressed iteration ranges (PR #211154)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 24 21:07:28 PDT 2026
https://github.com/qyingwu updated https://github.com/llvm/llvm-project/pull/211154
>From e8482a923674bdb6bdb92c65a2e1ba0a1ba3bc2c Mon Sep 17 00:00:00 2001
From: qyingwu <qiyingwu at utexas.edu>
Date: Tue, 21 Jul 2026 18:35:09 -0700
Subject: [PATCH] [mlir][sparse] Diagnose unsupported collapsed iteration
spaces
---
.../Transforms/SparseTensorPasses.cpp | 27 +++++++++++++++++++
.../Transforms/Utils/SparseTensorIterator.cpp | 7 +++++
.../SparseTensor/sparse_iteration_to_scf.mlir | 15 +++++++++++
.../sparse_iteration_to_scf_invalid.mlir | 15 +++++++++++
4 files changed, 64 insertions(+)
create mode 100644 mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf_invalid.mlir
diff --git a/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp
index b660e22154688..df8d696d75fd5 100644
--- a/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp
+++ b/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp
@@ -48,6 +48,28 @@ namespace {
// Passes implementation.
//===----------------------------------------------------------------------===//
+static bool supportsCollapsedRangeBetween(LevelType lt) {
+ return isDenseLT(lt) || isSingletonLT(lt);
+}
+
+static LogicalResult verifyLowerableCollapsedIterSpaces(Operation *op) {
+ WalkResult result = op->walk([&](ExtractIterSpaceOp op) {
+ auto [lvlLo, lvlHi] = op.getLvlRange();
+ SparseTensorEncodingAttr enc = cast<SparseTensorEncodingAttr>(
+ getRankedTensorType(op.getTensor()).getEncoding());
+ for (Level lvl = lvlLo + 1; lvl < lvlHi; ++lvl) {
+ LevelType lt = enc.getLvlType(lvl);
+ if (!supportsCollapsedRangeBetween(lt)) {
+ op.emitOpError() << "cannot lower collapsed iteration space with "
+ << toMLIRString(lt) << " level after the first level";
+ return WalkResult::interrupt();
+ }
+ }
+ return WalkResult::advance();
+ });
+ return failure(result.wasInterrupted());
+}
+
struct SparseAssembler : public impl::SparseAssemblerBase<SparseAssembler> {
SparseAssembler() = default;
SparseAssembler(const SparseAssembler &pass) = default;
@@ -167,6 +189,11 @@ struct LowerSparseIterationToSCFPass
default;
void runOnOperation() override {
+ if (failed(verifyLowerableCollapsedIterSpaces(getOperation()))) {
+ signalPassFailure();
+ return;
+ }
+
auto *ctx = &getContext();
RewritePatternSet patterns(ctx);
SparseIterationTypeConverter converter;
diff --git a/mlir/lib/Dialect/SparseTensor/Transforms/Utils/SparseTensorIterator.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/Utils/SparseTensorIterator.cpp
index b25e095d4c418..aa16eb837df3c 100644
--- a/mlir/lib/Dialect/SparseTensor/Transforms/Utils/SparseTensorIterator.cpp
+++ b/mlir/lib/Dialect/SparseTensor/Transforms/Utils/SparseTensorIterator.cpp
@@ -102,6 +102,12 @@ class DenseLevel : public SparseTensorLevel {
Value posLo = MULI(p, lvlSize);
return {posLo, lvlSize};
}
+
+ ValuePair collapseRangeBetween(OpBuilder &b, Location l, ValueRange,
+ ValuePair parentRange) const override {
+ return {MULI(parentRange.first, lvlSize),
+ MULI(parentRange.second, lvlSize)};
+ }
};
class BatchLevel : public SparseTensorLevel {
@@ -168,6 +174,7 @@ class CompressedLevel : public SparseLevel</*hasPosBuf=*/true> {
ValueRange posRange = posRangeIf.getResults();
return {posRange.front(), posRange.back()};
}
+
}; // namespace
class LooseCompressedLevel : public SparseLevel</*hasPosBuf=*/true> {
diff --git a/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf.mlir b/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf.mlir
index 855f1e99e7396..a649338bf02cc 100644
--- a/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf.mlir
+++ b/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf.mlir
@@ -55,6 +55,10 @@ func.func @sparse_iteration_to_scf(%sp : tensor<4x8xf32, #COO>) -> index {
map = (d0, d1) -> (d0 : dense, d1 : compressed)
}>
+#DenseDense = #sparse_tensor.encoding<{
+ map = (d0, d1) -> (d0 : dense, d1 : dense)
+}>
+
// CHECK-LABEL: @sparse_iteration_dense_level
// CHECK: scf.for
func.func @sparse_iteration_dense_level(%sp: tensor<?x?xf64, #DenseCompressed>) {
@@ -65,3 +69,14 @@ func.func @sparse_iteration_dense_level(%sp: tensor<?x?xf64, #DenseCompressed>)
}
return
}
+
+// CHECK-LABEL: @sparse_iteration_dense_dense_space
+// CHECK: scf.for
+func.func @sparse_iteration_dense_dense_space(%sp: tensor<?x?xf64, #DenseDense>) {
+ %0 = sparse_tensor.extract_iteration_space %sp lvls = 0 to 2
+ : tensor<?x?xf64, #DenseDense> -> !sparse_tensor.iter_space<#DenseDense, lvls = 0 to 2>
+ sparse_tensor.iterate %it in %0 at(%i, %j)
+ : !sparse_tensor.iter_space<#DenseDense, lvls = 0 to 2> {
+ }
+ return
+}
diff --git a/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf_invalid.mlir b/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf_invalid.mlir
new file mode 100644
index 0000000000000..316488c108eb1
--- /dev/null
+++ b/mlir/test/Dialect/SparseTensor/sparse_iteration_to_scf_invalid.mlir
@@ -0,0 +1,15 @@
+// RUN: mlir-opt %s --lower-sparse-iteration-to-scf -verify-diagnostics
+
+#DenseCompressed = #sparse_tensor.encoding<{
+ map = (d0, d1) -> (d0 : dense, d1 : compressed)
+}>
+
+func.func @unsupported_dense_compressed_space(%sp: tensor<?x?xf64, #DenseCompressed>) {
+ // expected-error at below {{cannot lower collapsed iteration space with compressed level after the first level}}
+ %0 = sparse_tensor.extract_iteration_space %sp lvls = 0 to 2
+ : tensor<?x?xf64, #DenseCompressed> -> !sparse_tensor.iter_space<#DenseCompressed, lvls = 0 to 2>
+ sparse_tensor.iterate %it in %0 at(%i, %j)
+ : !sparse_tensor.iter_space<#DenseCompressed, lvls = 0 to 2> {
+ }
+ return
+}
More information about the Mlir-commits
mailing list