[Mlir-commits] [mlir] [MLIR][Linalg] Fix FlattenElementwiseOp on broadcasted linalgs (PR #207005)

Hugo Trachino llvmlistbot at llvm.org
Tue Jul 7 08:08:32 PDT 2026


https://github.com/nujaa updated https://github.com/llvm/llvm-project/pull/207005

>From 3a79baaa2f3e2c39ee7479a7f716865626b88881 Mon Sep 17 00:00:00 2001
From: Hugo <hugo.trachino at huawei.com>
Date: Wed, 1 Jul 2026 21:08:28 +0800
Subject: [PATCH 1/2] [MLIR][Linalg] Fix FlattenElementwiseOp on broadcasted
 linalgs

---
 .../TransformOps/LinalgTransformOps.cpp       |  8 +++
 .../Dialect/Linalg/flatten-elementwise.mlir   | 53 +++++++++++++++++++
 2 files changed, 61 insertions(+)

diff --git a/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp b/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
index f44693096b26b..d39c1cb0e4a20 100644
--- a/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
+++ b/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
@@ -4297,6 +4297,14 @@ DiagnosedSilenceableFailure transform::FlattenElementwiseLinalgOp::applyToOne(
     return DiagnosedSilenceableFailure::success();
   }
 
+  // Bail out on linalgs with broadcasting semantics
+  if (!llvm::all_of(target.getIndexingMapsArray(), [](AffineMap m) {
+        return m.isProjectedPermutation(/*allowZeroInResults=*/false);
+      })) {
+    results.push_back(target);
+    return DiagnosedSilenceableFailure::success();
+  }
+
   // Attempt to flatten all dims to one.
   ReassociationIndices reassociation(target.getNumLoops());
   std::iota(reassociation.begin(), reassociation.end(), 0);
diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index 9fe50a521d2d8..1c6ba43888d7b 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -118,3 +118,56 @@ module attributes {transform.with_named_sequence} {
     transform.yield
   }
 }
+
+// -----
+// CHECK-LABEL: func.func @generic
+// CHECK-SAME:     %[[ARG0:[a-zA-Z0-9_]*]]: tensor<1x2x1xi32>
+// CHECK-SAME:     %[[ARG1:[a-zA-Z0-9_]*]]: tensor<32x2x2xi32>
+// CHECK-NEXT:     %[[RES:.*]] = linalg.generic
+// CHECK:          return %[[RES]]
+#map1 = affine_map<(d0, d1, d2) -> (0, d1, 0)>
+#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+func.func @generic(%arg0: tensor<1x2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
+    %13 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<1x2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
+      ^bb0(%in: i32, %out: i32):
+        linalg.yield %in : i32
+    } -> tensor<32x2x2xi32>
+    return %13 : tensor<32x2x2xi32>
+}
+
+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 @generic
+// CHECK-SAME:     %[[ARG0:[a-zA-Z0-9_]*]]: tensor<2x1xi32>
+// CHECK-SAME:     %[[ARG1:[a-zA-Z0-9_]*]]: tensor<32x2x2xi32>
+// CHECK-NEXT:     %[[RES:.*]] = linalg.generic
+// CHECK:          return %[[RES]]
+#map1 = affine_map<(d0, d1, d2) -> (d1, 0)>
+#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+func.func @generic(%arg0: tensor<2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
+    %13 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
+      ^bb0(%in: i32, %out: i32):
+        linalg.yield %in : i32
+    } -> tensor<32x2x2xi32>
+    return %13 : tensor<32x2x2xi32>
+}
+
+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
+  }
+}
\ No newline at end of file

>From c181d014aa9a96cd9e13d9cc9786f13eed40aff3 Mon Sep 17 00:00:00 2001
From: Hugo <hugo.trachino at huawei.com>
Date: Tue, 7 Jul 2026 23:07:32 +0800
Subject: [PATCH 2/2] fixup! [MLIR][Linalg] Fix FlattenElementwiseOp on
 broadcasted linalgs

---
 .../TransformOps/LinalgTransformOps.cpp       | 14 ++--
 .../Dialect/Linalg/flatten-elementwise.mlir   | 69 +++----------------
 .../Dialect/Linalg/flatten-unsupported.mlir   | 51 +++++++++++++-
 3 files changed, 63 insertions(+), 71 deletions(-)

diff --git a/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp b/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
index d39c1cb0e4a20..2b7d4e4ff6d8c 100644
--- a/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
+++ b/mlir/lib/Dialect/Linalg/TransformOps/LinalgTransformOps.cpp
@@ -4291,17 +4291,17 @@ DiagnosedSilenceableFailure transform::FlattenElementwiseLinalgOp::applyToOne(
     return mlir::emitSilenceableFailure(target->getLoc())
            << "only elementwise flattening is supported";
 
-  // If rank <= 1, do nothing
-  if (target.getNumLoops() <= 1) {
-    results.push_back(target);
-    return DiagnosedSilenceableFailure::success();
-  }
-
-  // Bail out on linalgs with broadcasting semantics
   if (!llvm::all_of(target.getIndexingMapsArray(), [](AffineMap m) {
         return m.isProjectedPermutation(/*allowZeroInResults=*/false);
       })) {
     results.push_back(target);
+    return mlir::emitSilenceableFailure(target->getLoc())
+           << "operators with broadcasting semantics are not supported";
+  }
+
+  // If rank <= 1, do nothing
+  if (target.getNumLoops() <= 1) {
+    results.push_back(target);
     return DiagnosedSilenceableFailure::success();
   }
 
diff --git a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
index 1c6ba43888d7b..f9f8f9e2b8b44 100644
--- a/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-elementwise.mlir
@@ -1,11 +1,11 @@
 // RUN: mlir-opt %s -transform-interpreter -split-input-file | FileCheck %s
 
-// CHECK-LABEL: func.func @fill(
+// CHECK-LABEL: func.func @fill_memref(
 // CHECK-SAME:                  %[[ARG0:.*]]: f32,
 // CHECK-SAME:                  %[[ARG1:.*]]: memref<32x7xf32>
 // CHECK-NEXT:    %[[FLATTENED:.*]] = memref.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
 // CHECK-NEXT:    linalg.fill ins(%[[ARG0]] : f32) outs(%[[FLATTENED]] : memref<224xf32>)
-func.func @fill(%cst: f32, %arg: memref<32x7xf32>) {
+func.func @fill_memref(%cst: f32, %arg: memref<32x7xf32>) {
     linalg.fill ins(%cst: f32) outs(%arg: memref<32x7xf32>)
     return
 }
@@ -43,7 +43,7 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
-// CHECK-LABEL: func.func @map(
+// CHECK-LABEL: func.func @map_memref(
 // CHECK-SAME:                 %[[ARG0:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG1:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG2:[a-zA-Z0-9_]*]]: memref<32x7xf32>
@@ -51,7 +51,7 @@ module attributes {transform.with_named_sequence} {
 // CHECK-NEXT:    %[[FLATTENED_1:.*]] = memref.collapse_shape %[[ARG1]] {{\[}}[0, 1]]
 // CHECK-NEXT:    %[[FLATTENED_2:.*]] = memref.collapse_shape %[[ARG2]] {{\[}}[0, 1]]
 // CHECK-NEXT:    linalg.map { arith.addf } ins(%[[FLATTENED_0]], %[[FLATTENED_1]] : memref<224xf32>, memref<224xf32>) outs(%[[FLATTENED_2]] : memref<224xf32>)
-func.func @map(%arg0: memref<32x7xf32>, %arg1: memref<32x7xf32>, %arg2: memref<32x7xf32>) {
+func.func @map_memref(%arg0: memref<32x7xf32>, %arg1: memref<32x7xf32>, %arg2: memref<32x7xf32>) {
     linalg.map {arith.addf} ins(%arg0, %arg1: memref<32x7xf32>, memref<32x7xf32>) outs(%arg2: memref<32x7xf32>)
     return
 }
@@ -67,12 +67,12 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
-// CHECK-LABEL: func.func @map_already_flat(
+// CHECK-LABEL: func.func @map_already_flat_memref(
 // CHECK-SAME:                 %[[ARG0:[a-zA-Z0-9_]*]]: memref<32xf32>
 // CHECK-SAME:                 %[[ARG1:[a-zA-Z0-9_]*]]: memref<32xf32>
 // CHECK-SAME:                 %[[ARG2:[a-zA-Z0-9_]*]]: memref<32xf32>
 // CHECK-NEXT:    linalg.map { arith.addf } ins(%[[ARG0]], %[[ARG1]] : memref<32xf32>, memref<32xf32>) outs(%[[ARG2]] : memref<32xf32>)
-func.func @map_already_flat(%arg0: memref<32xf32>, %arg1: memref<32xf32>, %arg2: memref<32xf32>) {
+func.func @map_already_flat_memref(%arg0: memref<32xf32>, %arg1: memref<32xf32>, %arg2: memref<32xf32>) {
     linalg.map {arith.addf} ins(%arg0, %arg1: memref<32xf32>, memref<32xf32>) outs(%arg2: memref<32xf32>)
     return
 }
@@ -89,7 +89,7 @@ module attributes {transform.with_named_sequence} {
 // -----
 
 // CHECK: #[[$MAP0:.*]] = affine_map<(d0) -> (d0)>
-// CHECK-LABEL: func.func @generic
+// CHECK-LABEL: func.func @elementwise_memref
 // CHECK-SAME:                 %[[ARG0:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG1:[a-zA-Z0-9_]*]]: memref<32x7xf32>
 // CHECK-SAME:                 %[[ARG2:[a-zA-Z0-9_]*]]: memref<32x7xf32>
@@ -101,7 +101,7 @@ module attributes {transform.with_named_sequence} {
 // CHECK-NEXT:         %[[SUM:.*]] = arith.addf %[[A]], %[[B]]
 // CHECK-NEXT:         linalg.yield %[[SUM]]
 #map = affine_map<(d0, d1) -> (d0, d1)>
-func.func @generic( %arg0: memref<32x7xf32>, %arg1: memref<32x7xf32>, %arg2: memref<32x7xf32>) {
+func.func @elementwise_memref( %arg0: memref<32x7xf32>, %arg1: memref<32x7xf32>, %arg2: memref<32x7xf32>) {
     linalg.generic {indexing_maps = [#map, #map, #map], iterator_types = ["parallel", "parallel"]} ins(%arg0, %arg1: memref<32x7xf32>, memref<32x7xf32>) outs(%arg2: memref<32x7xf32>) {
         ^bb0(%a: f32, %b: f32, %c: f32):
             %0 = arith.addf %a, %b : f32
@@ -118,56 +118,3 @@ module attributes {transform.with_named_sequence} {
     transform.yield
   }
 }
-
-// -----
-// CHECK-LABEL: func.func @generic
-// CHECK-SAME:     %[[ARG0:[a-zA-Z0-9_]*]]: tensor<1x2x1xi32>
-// CHECK-SAME:     %[[ARG1:[a-zA-Z0-9_]*]]: tensor<32x2x2xi32>
-// CHECK-NEXT:     %[[RES:.*]] = linalg.generic
-// CHECK:          return %[[RES]]
-#map1 = affine_map<(d0, d1, d2) -> (0, d1, 0)>
-#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-func.func @generic(%arg0: tensor<1x2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
-    %13 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<1x2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
-      ^bb0(%in: i32, %out: i32):
-        linalg.yield %in : i32
-    } -> tensor<32x2x2xi32>
-    return %13 : tensor<32x2x2xi32>
-}
-
-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 @generic
-// CHECK-SAME:     %[[ARG0:[a-zA-Z0-9_]*]]: tensor<2x1xi32>
-// CHECK-SAME:     %[[ARG1:[a-zA-Z0-9_]*]]: tensor<32x2x2xi32>
-// CHECK-NEXT:     %[[RES:.*]] = linalg.generic
-// CHECK:          return %[[RES]]
-#map1 = affine_map<(d0, d1, d2) -> (d1, 0)>
-#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-func.func @generic(%arg0: tensor<2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
-    %13 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
-      ^bb0(%in: i32, %out: i32):
-        linalg.yield %in : i32
-    } -> tensor<32x2x2xi32>
-    return %13 : tensor<32x2x2xi32>
-}
-
-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
-  }
-}
\ No newline at end of file
diff --git a/mlir/test/Dialect/Linalg/flatten-unsupported.mlir b/mlir/test/Dialect/Linalg/flatten-unsupported.mlir
index 499db4cfb3294..f0b8b8b153bbb 100644
--- a/mlir/test/Dialect/Linalg/flatten-unsupported.mlir
+++ b/mlir/test/Dialect/Linalg/flatten-unsupported.mlir
@@ -1,6 +1,6 @@
 // RUN: mlir-opt %s -transform-interpreter -split-input-file -verify-diagnostics
 
-func.func @non_elementwise(%arg0: memref<2x3xf32>, %arg1: memref<3x4xf32>, %arg2: memref<2x4xf32>) {
+func.func @unsupported_non_elementwise(%arg0: memref<2x3xf32>, %arg1: memref<3x4xf32>, %arg2: memref<2x4xf32>) {
   // expected-error @below {{only elementwise flattening is supported}}
   linalg.matmul ins(%arg0, %arg1 : memref<2x3xf32>, memref<3x4xf32>) outs(%arg2: memref<2x4xf32>)
   return
@@ -19,8 +19,53 @@ module attributes {transform.with_named_sequence} {
 
 func.func @unsupported_memref(%arg0: memref<32x7xf32, strided<[7, 2]>>, %arg1: memref<32x7xf32, strided<[7, 2]>>, %arg2: memref<32x7xf32, strided<[7, 2]>>) {
   // expected-error @below {{attempted to flatten, but failed}}
-    linalg.map {arith.addf} ins(%arg0, %arg1: memref<32x7xf32, strided<[7, 2]>>, memref<32x7xf32, strided<[7, 2]>>) outs(%arg2: memref<32x7xf32, strided<[7, 2]>>)
-    return
+  linalg.map {arith.addf} ins(%arg0, %arg1: memref<32x7xf32, strided<[7, 2]>>, memref<32x7xf32, strided<[7, 2]>>) outs(%arg2: memref<32x7xf32, strided<[7, 2]>>)
+  return
+}
+
+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
+  }
+}
+
+// -----
+#map1 = affine_map<(d0, d1, d2) -> (0, d1, 0)>
+#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+func.func @unsupported_broadcasting_elementwise(%arg0: tensor<1x2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
+  // expected-error @below {{operators with broadcasting semantics are not supported}}
+  %0 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<1x2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
+    ^bb0(%in: i32, %out: i32):
+      linalg.yield %in : i32
+  } -> tensor<32x2x2xi32>
+  return %0 : tensor<32x2x2xi32>
+}
+
+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
+  }
+}
+
+// -----
+
+#map1 = affine_map<(d0, d1, d2) -> (d1, 0)>
+#map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+func.func @unsupported_rank_expanding_broadcasting_elementwise(%arg0: tensor<2x1xi32>, %arg1: tensor<32x2x2xi32>) -> tensor<32x2x2xi32> {
+  // expected-error @below {{operators with broadcasting semantics are not supported}}
+  %0 = linalg.generic {indexing_maps = [#map1, #map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%arg0 : tensor<2x1xi32>) outs(%arg1 : tensor<32x2x2xi32>) {
+    ^bb0(%in: i32, %out: i32):
+      linalg.yield %in : i32
+  } -> tensor<32x2x2xi32>
+  return %0 : tensor<32x2x2xi32>
 }
 
 module attributes {transform.with_named_sequence} {



More information about the Mlir-commits mailing list