[Mlir-commits] [mlir] 76525a8 - [MLIR] Fix canonicalization of extract_slice(unpack) (#181840)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu Feb 19 13:12:29 PST 2026
Author: Mikhail Romanov
Date: 2026-02-19T21:12:24Z
New Revision: 76525a8d3b5caeb849635a1fe3438dcde315dc2c
URL: https://github.com/llvm/llvm-project/commit/76525a8d3b5caeb849635a1fe3438dcde315dc2c
DIFF: https://github.com/llvm/llvm-project/commit/76525a8d3b5caeb849635a1fe3438dcde315dc2c.diff
LOG: [MLIR] Fix canonicalization of extract_slice(unpack) (#181840)
Added:
Modified:
mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
mlir/test/Dialect/Linalg/canonicalize.mlir
Removed:
################################################################################
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
index 161bdaaf5bc6c..685096de5fbea 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
@@ -6399,8 +6399,10 @@ bool UnPackOp::canFoldSliceOp(tensor::ExtractSliceOp sliceOp) {
RankedTensorType unpackedTypeAfterFold = sliceOp.getResultType();
SmallVector<int64_t> outerShapeWithoutTranspose =
getPackedOuterShapeWithoutTransposition(*this);
+ SmallVector<bool> areOuterDimsTiled(outerShapeWithoutTranspose.size(), false);
for (auto [pos, tileSize] :
llvm::zip_equal(this->getInnerDimsPos(), this->getStaticInnerTiles())) {
+ areOuterDimsTiled[pos] = true;
if (unpackedTypeAfterFold.isDynamicDim(pos))
return false;
if (ShapedType::isDynamic(outerShapeWithoutTranspose[pos]))
@@ -6412,6 +6414,16 @@ bool UnPackOp::canFoldSliceOp(tensor::ExtractSliceOp sliceOp) {
if (paddingSize >= tileSize)
return false;
}
+ // extract_slice must not affect dimensions that are not being unpacked
+ for (int64_t pos = 0, e = outerShapeWithoutTranspose.size(); pos < e; ++pos) {
+ if (areOuterDimsTiled[pos])
+ continue;
+ int64_t dim = outerShapeWithoutTranspose[pos];
+ if (ShapedType::isDynamic(dim))
+ return false;
+ if (dim != unpackedTypeAfterFold.getDimSize(pos))
+ return false;
+ }
return true;
}
diff --git a/mlir/test/Dialect/Linalg/canonicalize.mlir b/mlir/test/Dialect/Linalg/canonicalize.mlir
index 08eb40d4cb442..77c1c3da17166 100644
--- a/mlir/test/Dialect/Linalg/canonicalize.mlir
+++ b/mlir/test/Dialect/Linalg/canonicalize.mlir
@@ -2061,6 +2061,52 @@ func.func @no_fold_extract_slice_into_unpack_non_zero_offset(
// -----
+// Must not fold because extract_slice cuts the 0'th dimension from 30 to 28.
+func.func @no_fold_extract_slice_into_unpack_slice_over_non_tiled_dim(
+ %src : tensor<30x2x16xf32>, %dest : tensor<30x32xf32>
+) -> tensor<28x28xf32> {
+ %unpack = linalg.unpack %src
+ inner_dims_pos = [1]
+ inner_tiles = [16]
+ into %dest : tensor<30x2x16xf32> -> tensor<30x32xf32>
+ %extracted_slice = tensor.extract_slice %unpack
+ [0, 0] [28, 28] [1, 1] : tensor<30x32xf32> to tensor<28x28xf32>
+ return %extracted_slice : tensor<28x28xf32>
+}
+
+// CHECK-LABEL: func @no_fold_extract_slice_into_unpack_slice_over_non_tiled_dim
+// CHECK-SAME: %[[SRC:.+]]: tensor<30x2x16xf32>
+// CHECK-SAME: %[[DEST:.+]]: tensor<30x32xf32>
+// CHECK: %[[UNPACK:.+]] = linalg.unpack %[[SRC]]
+// CHECK-SAME: into %[[DEST]]
+// CHECK: %[[SLICE:.+]] = tensor.extract_slice %[[UNPACK]]
+// CHECK: return %[[SLICE]]
+
+// -----
+
+// Must not fold because extract_slice's effect on the 0'th dimension is unknown.
+func.func @no_fold_extract_slice_into_unpack_slice_over_dynamic_dim(
+ %src : tensor<?x2x16xf32>, %dest : tensor<?x32xf32>, %size : index
+) -> tensor<?x28xf32> {
+ %unpack = linalg.unpack %src
+ inner_dims_pos = [1]
+ inner_tiles = [16]
+ into %dest : tensor<?x2x16xf32> -> tensor<?x32xf32>
+ %extracted_slice = tensor.extract_slice %unpack
+ [0, 0] [%size, 28] [1, 1] : tensor<?x32xf32> to tensor<?x28xf32>
+ return %extracted_slice : tensor<?x28xf32>
+}
+
+// CHECK-LABEL: func @no_fold_extract_slice_into_unpack_slice_over_dynamic_dim
+// CHECK-SAME: %[[SRC:.+]]: tensor<?x2x16xf32>
+// CHECK-SAME: %[[DEST:.+]]: tensor<?x32xf32>
+// CHECK: %[[UNPACK:.+]] = linalg.unpack %[[SRC]]
+// CHECK-SAME: into %[[DEST]]
+// CHECK: %[[SLICE:.+]] = tensor.extract_slice %[[UNPACK]]
+// CHECK: return %[[SLICE]]
+
+// -----
+
// CHECK-LABEL: func.func @fold_cast_unpack_dynamic_tile_size(
// CHECK-SAME: %[[SRC:.*]]: tensor<1x1x8x1xi32>,
// CHECK-SAME: %[[DEST:.*]]: tensor<7x?xi32>) -> tensor<7x?xi32> {
More information about the Mlir-commits
mailing list