[Mlir-commits] [mlir] [MLIR][Linalg] Recompute linalg.broadcast dimensions when flattening (PR #213641)

Chibuoyim Ogbonna llvmlistbot at llvm.org
Mon Aug 10 03:59:59 PDT 2026


https://github.com/bruteforceboy updated https://github.com/llvm/llvm-project/pull/213641

>From 7111d16d6cf796d3766eada85498f3292f879e61 Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Mon, 3 Aug 2026 22:53:55 +0800
Subject: [PATCH 1/6] [MLIR][Linalg] Recompute linalg.broadcast dimensions when
 flattening

---
 .../Linalg/Transforms/ElementwiseOpFusion.cpp | 40 +++++++++++++++++
 .../Dialect/Linalg/flatten-elementwise.mlir   | 45 +++++++++++++++++++
 2 files changed, 85 insertions(+)

diff --git a/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp b/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
index db46de75abd1a..5282514f9c5dc 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
@@ -1809,12 +1809,49 @@ GenericOp cloneToCollapsedOp<GenericOp>(RewriterBase &rewriter,
   return collapsedOp;
 }
 
+/// Collapse a `BroadcastOp`, recomputing its `dimensions` for the collapsed
+/// iteration space. Returns null if the collapse is not expressible as a
+/// broadcast.
+template <>
+BroadcastOp
+cloneToCollapsedOp<BroadcastOp>(RewriterBase &rewriter, BroadcastOp origOp,
+                                const CollapsingInfo &collapsingInfo) {
+  ArrayRef<int64_t> broadcastDims = origOp.getDimensions();
+  SmallVector<int64_t> newDimensions;
+  for (auto [collapsedDim, foldedDims] :
+       llvm::enumerate(collapsingInfo.getCollapsedOpToOrigOpMapping())) {
+    size_t numBroadcast = llvm::count_if(foldedDims, [&](int64_t d) {
+      return llvm::is_contained(broadcastDims, d);
+    });
+    // A collapsed dimension is a broadcast dimension iff all the dimensions it
+    // folds are; a mix of the two cannot be represented as a broadcast.
+    if (numBroadcast != 0 && numBroadcast != foldedDims.size())
+      return nullptr;
+    if (numBroadcast != 0)
+      newDimensions.push_back(collapsedDim);
+  }
+
+  SmallVector<Value> inputOperands, outputOperands;
+  SmallVector<Type> resultTypes;
+  collapseOperandsAndResults(origOp, collapsingInfo, rewriter, inputOperands,
+                             outputOperands, resultTypes);
+
+  return BroadcastOp::create(rewriter, origOp.getLoc(), inputOperands[0],
+                             outputOperands[0], newDimensions);
+}
+
 static LinalgOp createCollapsedOp(LinalgOp op,
                                   const CollapsingInfo &collapsingInfo,
                                   RewriterBase &rewriter) {
   if (GenericOp genericOp = dyn_cast<GenericOp>(op.getOperation())) {
     return cloneToCollapsedOp(rewriter, genericOp, collapsingInfo);
   }
+  if (BroadcastOp broadcastOp = dyn_cast<BroadcastOp>(op.getOperation())) {
+    BroadcastOp collapsedOp =
+        cloneToCollapsedOp(rewriter, broadcastOp, collapsingInfo);
+    return collapsedOp ? cast<LinalgOp>(collapsedOp.getOperation())
+                       : LinalgOp();
+  }
   return cloneToCollapsedOp(rewriter, op, collapsingInfo);
 }
 
@@ -1871,6 +1908,9 @@ FailureOr<CollapseResult> mlir::linalg::collapseOpIterationDims(
   }
 
   LinalgOp collapsedOp = createCollapsedOp(op, collapsingInfo, rewriter);
+  if (!collapsedOp)
+    return rewriter.notifyMatchFailure(
+        op, "failed to create collapsed op for the specified dimensions");
 
   Location loc = op->getLoc();
   SmallVector<OpFoldResult> loopBound =
diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index ca06062f61840..451563d25ca1c 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -71,6 +71,51 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
+// CHECK-LABEL: func.func @broadcast_rank0_named_tensor(
+// CHECK-SAME:                         %[[ARG0:.*]]: tensor<i32>,
+// CHECK-SAME:                         %[[ARG1:.*]]: tensor<32x2xi32>
+// CHECK-NEXT:    %[[FLATTENED:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
+// CHECK-NEXT:    %[[FLATTENED_RESULT:.*]] = linalg.broadcast ins(%[[ARG0]] : tensor<i32>) outs(%[[FLATTENED]] : tensor<64xi32>) dimensions = [0]
+// CHECK:         %[[RESULT:.*]] = tensor.expand_shape %[[FLATTENED_RESULT]] {{\[}}[0, 1]] output_shape [32, 2] : tensor<64xi32> into tensor<32x2xi32>
+func.func @broadcast_rank0_named_tensor(%arg0: tensor<i32>, %arg1: tensor<32x2xi32>) -> tensor<32x2xi32> {
+  %0 = linalg.broadcast ins(%arg0 : tensor<i32>) outs(%arg1 : tensor<32x2xi32>) dimensions = [0, 1]
+  return %0 : tensor<32x2xi32>
+}
+
+module attributes {transform.with_named_sequence} {
+  transform.named_sequence @__transform_main(%arg1: !transform.any_op {transform.readonly}) {
+    %0 = transform.structured.match interface{LinalgOp} in %arg1 : (!transform.any_op) -> !transform.any_op
+    %flattened = transform.structured.flatten_elementwise %0
+      : (!transform.any_op) -> !transform.any_op
+    transform.yield
+  }
+}
+
+// -----
+
+// CHECK-LABEL: func.func @broadcast_identity_named_tensor(
+// CHECK-SAME:                         %[[ARG0:.*]]: tensor<4x8xf32>,
+// CHECK-SAME:                         %[[ARG1:.*]]: tensor<4x8xf32>
+// CHECK-NEXT:    %[[IN:.*]] = tensor.collapse_shape %[[ARG0]] {{\[}}[0, 1]]
+// CHECK-NEXT:    %[[OUT:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
+// CHECK-NEXT:    %[[FLATTENED_RESULT:.*]] = linalg.broadcast ins(%[[IN]] : tensor<32xf32>) outs(%[[OUT]] : tensor<32xf32>) dimensions = []
+// CHECK:         %[[RESULT:.*]] = tensor.expand_shape %[[FLATTENED_RESULT]] {{\[}}[0, 1]] output_shape [4, 8] : tensor<32xf32> into tensor<4x8xf32>
+func.func @broadcast_identity_named_tensor(%arg0: tensor<4x8xf32>, %arg1: tensor<4x8xf32>) -> tensor<4x8xf32> {
+  %0 = linalg.broadcast ins(%arg0 : tensor<4x8xf32>) outs(%arg1 : tensor<4x8xf32>) dimensions = []
+  return %0 : tensor<4x8xf32>
+}
+
+module attributes {transform.with_named_sequence} {
+  transform.named_sequence @__transform_main(%arg1: !transform.any_op {transform.readonly}) {
+    %0 = transform.structured.match interface{LinalgOp} in %arg1 : (!transform.any_op) -> !transform.any_op
+    %flattened = transform.structured.flatten_elementwise %0
+      : (!transform.any_op) -> !transform.any_op
+    transform.yield
+  }
+}
+
+// -----
+
 // CHECK-LABEL: func.func @map_memref(
 // CHECK-SAME:                 %[[ARG0:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG1:[a-zA-Z0-9_]*]]: memref<32x7xf32>

>From 8c4527a3f9ca1c30e0ce29c5f134775f58a27ac2 Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Fri, 7 Aug 2026 18:24:56 +0800
Subject: [PATCH 2/6] Apply review comments: Simplify collapse logic

---
 .../Linalg/Transforms/ElementwiseOpFusion.cpp | 32 ++++---------------
 1 file changed, 7 insertions(+), 25 deletions(-)

diff --git a/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp b/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
index 5282514f9c5dc..2d367d72ae6e1 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/ElementwiseOpFusion.cpp
@@ -1809,33 +1809,21 @@ GenericOp cloneToCollapsedOp<GenericOp>(RewriterBase &rewriter,
   return collapsedOp;
 }
 
-/// Collapse a `BroadcastOp`, recomputing its `dimensions` for the collapsed
-/// iteration space. Returns null if the collapse is not expressible as a
-/// broadcast.
+/// Collapse a `BroadcastOp`. Flattening leaves a single dimension, so a 0-D
+/// input broadcasts into it (`dimensions = [0]`) and any other input adds none.
 template <>
 BroadcastOp
 cloneToCollapsedOp<BroadcastOp>(RewriterBase &rewriter, BroadcastOp origOp,
                                 const CollapsingInfo &collapsingInfo) {
-  ArrayRef<int64_t> broadcastDims = origOp.getDimensions();
-  SmallVector<int64_t> newDimensions;
-  for (auto [collapsedDim, foldedDims] :
-       llvm::enumerate(collapsingInfo.getCollapsedOpToOrigOpMapping())) {
-    size_t numBroadcast = llvm::count_if(foldedDims, [&](int64_t d) {
-      return llvm::is_contained(broadcastDims, d);
-    });
-    // A collapsed dimension is a broadcast dimension iff all the dimensions it
-    // folds are; a mix of the two cannot be represented as a broadcast.
-    if (numBroadcast != 0 && numBroadcast != foldedDims.size())
-      return nullptr;
-    if (numBroadcast != 0)
-      newDimensions.push_back(collapsedDim);
-  }
-
   SmallVector<Value> inputOperands, outputOperands;
   SmallVector<Type> resultTypes;
   collapseOperandsAndResults(origOp, collapsingInfo, rewriter, inputOperands,
                              outputOperands, resultTypes);
 
+  SmallVector<int64_t> newDimensions;
+  if (origOp.getInput().getType().getRank() == 0)
+    newDimensions.push_back(0);
+
   return BroadcastOp::create(rewriter, origOp.getLoc(), inputOperands[0],
                              outputOperands[0], newDimensions);
 }
@@ -1847,10 +1835,7 @@ static LinalgOp createCollapsedOp(LinalgOp op,
     return cloneToCollapsedOp(rewriter, genericOp, collapsingInfo);
   }
   if (BroadcastOp broadcastOp = dyn_cast<BroadcastOp>(op.getOperation())) {
-    BroadcastOp collapsedOp =
-        cloneToCollapsedOp(rewriter, broadcastOp, collapsingInfo);
-    return collapsedOp ? cast<LinalgOp>(collapsedOp.getOperation())
-                       : LinalgOp();
+    return cloneToCollapsedOp(rewriter, broadcastOp, collapsingInfo);
   }
   return cloneToCollapsedOp(rewriter, op, collapsingInfo);
 }
@@ -1908,9 +1893,6 @@ FailureOr<CollapseResult> mlir::linalg::collapseOpIterationDims(
   }
 
   LinalgOp collapsedOp = createCollapsedOp(op, collapsingInfo, rewriter);
-  if (!collapsedOp)
-    return rewriter.notifyMatchFailure(
-        op, "failed to create collapsed op for the specified dimensions");
 
   Location loc = op->getLoc();
   SmallVector<OpFoldResult> loopBound =

>From 03a54bb233c63830c20873019ad5c65061fe76d9 Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Fri, 7 Aug 2026 18:28:17 +0800
Subject: [PATCH 3/6] Apply review comments: remove identity lit test

---
 .../Dialect/Linalg/flatten-elementwise.mlir   | 23 -------------------
 1 file changed, 23 deletions(-)

diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index 451563d25ca1c..4c8e61835a1aa 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -93,29 +93,6 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
-// CHECK-LABEL: func.func @broadcast_identity_named_tensor(
-// CHECK-SAME:                         %[[ARG0:.*]]: tensor<4x8xf32>,
-// CHECK-SAME:                         %[[ARG1:.*]]: tensor<4x8xf32>
-// CHECK-NEXT:    %[[IN:.*]] = tensor.collapse_shape %[[ARG0]] {{\[}}[0, 1]]
-// CHECK-NEXT:    %[[OUT:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
-// CHECK-NEXT:    %[[FLATTENED_RESULT:.*]] = linalg.broadcast ins(%[[IN]] : tensor<32xf32>) outs(%[[OUT]] : tensor<32xf32>) dimensions = []
-// CHECK:         %[[RESULT:.*]] = tensor.expand_shape %[[FLATTENED_RESULT]] {{\[}}[0, 1]] output_shape [4, 8] : tensor<32xf32> into tensor<4x8xf32>
-func.func @broadcast_identity_named_tensor(%arg0: tensor<4x8xf32>, %arg1: tensor<4x8xf32>) -> tensor<4x8xf32> {
-  %0 = linalg.broadcast ins(%arg0 : tensor<4x8xf32>) outs(%arg1 : tensor<4x8xf32>) dimensions = []
-  return %0 : tensor<4x8xf32>
-}
-
-module attributes {transform.with_named_sequence} {
-  transform.named_sequence @__transform_main(%arg1: !transform.any_op {transform.readonly}) {
-    %0 = transform.structured.match interface{LinalgOp} in %arg1 : (!transform.any_op) -> !transform.any_op
-    %flattened = transform.structured.flatten_elementwise %0
-      : (!transform.any_op) -> !transform.any_op
-    transform.yield
-  }
-}
-
-// -----
-
 // CHECK-LABEL: func.func @map_memref(
 // CHECK-SAME:                 %[[ARG0:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG1:[a-zA-Z0-9_]*]]: memref<32x7xf32>

>From 9630040045a1e39c6b0f1a36f4479ffe7a28ede8 Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Mon, 10 Aug 2026 11:37:10 +0100
Subject: [PATCH 4/6] Update mlir/test/Dialect/Linalg/flatten-elementwise.mlir
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

Co-authored-by: Andrzej Warzyński <andrzej.warzynski at gmail.com>
---
 mlir/test/Dialect/Linalg/flatten-elementwise.mlir | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index 4c8e61835a1aa..f88a594f0e621 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -52,7 +52,7 @@ module attributes {transform.with_named_sequence} {
 #map0 = affine_map<(d0, d1) -> ()>
 #map1 = affine_map<(d0, d1) -> (d0, d1)>
 
-func.func @broadcast_rank0_tensor(%arg0: tensor<i32>, %arg1: tensor<32x2xi32>) -> tensor<32x2xi32> {
+func.func @broadcast_as_generic_rank0_tensor(%arg0: tensor<i32>, %arg1: tensor<32x2xi32>) -> tensor<32x2xi32> {
   %0 = linalg.generic {indexing_maps = [#map0, #map1], iterator_types = ["parallel", "parallel"]} ins(%arg0 : tensor<i32>) outs(%arg1 : tensor<32x2xi32>) {
     ^bb0(%in: i32, %out: i32):
       linalg.yield %in : i32

>From 655602e8077077e2b18aaf05b1e03c9ef2aa9ac7 Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Mon, 10 Aug 2026 11:37:20 +0100
Subject: [PATCH 5/6] Update mlir/test/Dialect/Linalg/flatten-elementwise.mlir
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

Co-authored-by: Andrzej Warzyński <andrzej.warzynski at gmail.com>
---
 mlir/test/Dialect/Linalg/flatten-elementwise.mlir | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index f88a594f0e621..173c0d9482063 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -77,7 +77,7 @@ module attributes {transform.with_named_sequence} {
 // CHECK-NEXT:    %[[FLATTENED:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
 // CHECK-NEXT:    %[[FLATTENED_RESULT:.*]] = linalg.broadcast ins(%[[ARG0]] : tensor<i32>) outs(%[[FLATTENED]] : tensor<64xi32>) dimensions = [0]
 // CHECK:         %[[RESULT:.*]] = tensor.expand_shape %[[FLATTENED_RESULT]] {{\[}}[0, 1]] output_shape [32, 2] : tensor<64xi32> into tensor<32x2xi32>
-func.func @broadcast_rank0_named_tensor(%arg0: tensor<i32>, %arg1: tensor<32x2xi32>) -> tensor<32x2xi32> {
+func.func @broadcast_as_named_rank0_tensor(%arg0: tensor<i32>, %arg1: tensor<32x2xi32>) -> tensor<32x2xi32> {
   %0 = linalg.broadcast ins(%arg0 : tensor<i32>) outs(%arg1 : tensor<32x2xi32>) dimensions = [0, 1]
   return %0 : tensor<32x2xi32>
 }

>From 64a9198223b0d3f3ee1144e81c8e980ae724b5be Mon Sep 17 00:00:00 2001
From: bruteforceboy <chibuoyim.faith.ogbonna at huawei.com>
Date: Mon, 10 Aug 2026 18:56:58 +0800
Subject: [PATCH 6/6] Apply review comment: update flatten-elementwise.mlir

---
 mlir/test/Dialect/Linalg/flatten-elementwise.mlir | 4 ++--
 1 file changed, 2 insertions(+), 2 deletions(-)

diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index 173c0d9482063..9eb0131c9edd1 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -43,7 +43,7 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
-// CHECK-LABEL: func.func @broadcast_rank0_tensor(
+// CHECK-LABEL: func.func @broadcast_as_generic_rank0_tensor(
 // CHECK-SAME:                         %[[ARG0:.*]]: tensor<i32>,
 // CHECK-SAME:                         %[[ARG1:.*]]: tensor<32x2xi32>
 // CHECK-NEXT:    %[[FLATTENED:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
@@ -71,7 +71,7 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
-// CHECK-LABEL: func.func @broadcast_rank0_named_tensor(
+// CHECK-LABEL: func.func @broadcast_as_named_rank0_tensor(
 // CHECK-SAME:                         %[[ARG0:.*]]: tensor<i32>,
 // CHECK-SAME:                         %[[ARG1:.*]]: tensor<32x2xi32>
 // CHECK-NEXT:    %[[FLATTENED:.*]] = tensor.collapse_shape %[[ARG1]] {{\[}}[0, 1]]



More information about the Mlir-commits mailing list