[llvm-branch-commits] [mlir] [MLIR][Linalg] Remove ternary named op (PR #220916)

Renato Golin via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Fri Sep 4 09:27:31 PDT 2026


https://github.com/rengolin updated https://github.com/llvm/llvm-project/pull/220916

>From 784116c3999d3b9938038a604e837855e7341042 Mon Sep 17 00:00:00 2001
From: Renato Golin <rengolin at systemcall.eu>
Date: Tue, 11 Aug 2026 17:11:26 +0100
Subject: [PATCH 1/3] [MLIR][Linalg] Remove ternary named op

Removes the named op `select` from the Linalg dialect.

I have also updated the ElementwiseOp builder to simplify the default
case: kind + no affine map.

Depends on both unary and binary removal branches.

Ref:
https://discourse.llvm.org/t/rfc-update-semantics-of-linalg-named-operations-unary-binary-ternary/91531
---
 .../Linalg/IR/LinalgNamedStructuredOps.yaml   | 57 -----------------
 .../Dialect/Linalg/IR/LinalgStructuredOps.td  | 16 ++++-
 .../Dialect/Linalg/Transforms/Transforms.h    |  9 ---
 .../Conversion/TosaToLinalg/TosaToLinalg.cpp  |  3 +-
 .../Dialect/Linalg/Transforms/CMakeLists.txt  |  2 -
 .../Linalg/Transforms/CategoryToNamedOp.cpp   | 61 -------------------
 .../Dialect/Linalg/Transforms/MorphOps.cpp    |  4 --
 .../Linalg/Transforms/NamedToElementwise.cpp  | 57 -----------------
 .../Dialect/Linalg/Transforms/Specialize.cpp  |  3 +-
 .../linalg/opdsl/ops/core_named_ops.py        | 20 ------
 .../elementwise/named-to-elementwise.mlir     | 17 ------
 .../Dialect/Linalg/generalize-named-ops.mlir  | 26 --------
 .../linalg-morph-elementwise-to-named.mlir    | 18 ------
 mlir/test/Dialect/Linalg/named-ops-fail.mlir  | 48 ---------------
 mlir/test/Dialect/Linalg/named-ops.mlir       | 48 ---------------
 15 files changed, 16 insertions(+), 373 deletions(-)
 delete mode 100644 mlir/lib/Dialect/Linalg/Transforms/CategoryToNamedOp.cpp
 delete mode 100644 mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
 delete mode 100644 mlir/test/Dialect/Linalg/elementwise/named-to-elementwise.mlir
 delete mode 100644 mlir/test/Dialect/Linalg/named-ops-fail.mlir

diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
index 58ef08fee0463..828981fe17a3f 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
@@ -44,63 +44,6 @@ structured_op: !LinalgStructuredOpConfig
         - !ScalarExpression
           scalar_arg: I
 --- !LinalgOpConfig
-metadata: !LinalgOpMetadata
-  name: select
-  cpp_class_name: SelectOp
-  doc: |-
-    Chooses one value based on a binary condition supplied as its first operand.
-
-    The shapes and element types must be identical. The appropriate casts,
-    broadcasts and reductions should be done previously to calling this op.
-
-    This means reduction/broadcast/element cast semantics is explicit. Further
-    passes can take that into account when lowering this code. For example,
-    a `linalg.broadcast` + `linalg.select` sequence can be lowered to a
-    `linalg.generic` with different affine maps for the two operands.
-structured_op: !LinalgStructuredOpConfig
-  args:
-  - !LinalgOperandDefConfig
-    name: cond
-    kind: input_tensor
-    type_var: U
-    shape_map: affine_map<() -> ()>
-  - !LinalgOperandDefConfig
-    name: lhs
-    kind: input_tensor
-    type_var: T1
-    shape_map: affine_map<() -> ()>
-  - !LinalgOperandDefConfig
-    name: rhs
-    kind: input_tensor
-    type_var: T1
-    shape_map: affine_map<() -> ()>
-  - !LinalgOperandDefConfig
-    name: O
-    kind: output_tensor
-    type_var: T1
-    shape_map: affine_map<() -> ()>
-  indexing_maps: !LinalgIndexingMapsConfig
-    static_indexing_maps:
-    - affine_map<() -> ()>
-    - affine_map<() -> ()>
-    - affine_map<() -> ()>
-    - affine_map<() -> ()>
-  iterator_types: []
-  assignments:
-  - !ScalarAssign
-    arg: O
-    value: !ScalarExpression
-      scalar_fn:
-        kind: ternary
-        fn_name: select
-        operands:
-        - !ScalarExpression
-          scalar_arg: cond
-        - !ScalarExpression
-          scalar_arg: lhs
-        - !ScalarExpression
-          scalar_arg: rhs
---- !LinalgOpConfig
 metadata: !LinalgOpMetadata
   name: quantized_matmul
   cpp_class_name: QuantizedMatmulOp
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td b/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
index eb240fa63c4ed..5e07d3b8e551c 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
@@ -610,11 +610,21 @@ def ElementwiseOp : LinalgStructuredBase_Op<"elementwise", [
       }]>,
 
      OpBuilder<(ins "ValueRange":$inputs, "ValueRange":$outputs,
-          "ElementwiseKindAttr":$kind,
-          "ArrayAttr":$indexingMaps,
+          "ElementwiseKind":$kind,
+          CArg<"ArrayAttr", "{}">:$indexingMaps,
           CArg<"ArrayRef<NamedAttribute>", "{}">:$attributes),
       [{
-        $_state.addAttribute("kind", kind);
+        assert((unsigned)kind <= getMaxEnumValForElementwiseKind() &&
+               "expected a valid elementwise kind attribute");
+        ElementwiseKindAttr kindAttr = ElementwiseKindAttr::get($_builder.getContext(), kind);
+        $_state.addAttribute("kind", kindAttr);
+        if (!indexingMaps) {
+          auto affineMaps = ElementwiseOp::getDefaultIndexingMaps(
+                            inputs.size() + outputs.size(),
+                            llvm::cast<ShapedType>(outputs[0].getType()).getRank(),
+                            $_builder.getContext());
+          indexingMaps = $_builder.getAffineMapArrayAttr(affineMaps);
+        }
         $_state.addAttribute("indexing_maps", indexingMaps);
         buildStructuredOp($_builder, $_state, std::nullopt, inputs, outputs,
                           attributes, ElementwiseOp::getRegionBuilder());
diff --git a/mlir/include/mlir/Dialect/Linalg/Transforms/Transforms.h b/mlir/include/mlir/Dialect/Linalg/Transforms/Transforms.h
index 9ab50f1136ac1..cbc338c9f1b32 100644
--- a/mlir/include/mlir/Dialect/Linalg/Transforms/Transforms.h
+++ b/mlir/include/mlir/Dialect/Linalg/Transforms/Transforms.h
@@ -1923,15 +1923,6 @@ void populateLinalgGenericOpsSpecializationPatterns(
     RewritePatternSet &patterns,
     const GenericOpSpecializationOptions &options = {});
 
-/// Populates `patterns` that convert linalg named ops e.g. `linalg.add`
-/// to equivalent `linalg.elementwise`.
-void populateLinalgNamedToElementwisePatterns(RewritePatternSet &patterns);
-
-/// Populates `patterns` that convert linalg category ops (e.g.
-/// `linalg.elementwise`, `linalg.contract`) to equivalent linalg named ops
-/// (e.g. `linalg.add`, `linalg.matmul`).
-void populateLinalgCategoryToNamedPatterns(RewritePatternSet &patterns);
-
 /// Populates `patterns` with patterns that fold operations like
 /// `linalg.transform` into elementwise op map.
 void populateLinalgFoldIntoElementwisePatterns(RewritePatternSet &patterns);
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index b7eb0a3aed546..cebf40bbe2c38 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -1321,7 +1321,8 @@ static LogicalResult reduceMatchAndRewriteHelper(OpTy op, uint64_t axis,
     ins.push_back(linalgOp->getResult(0));
     outs.push_back(finalEmptyTensor);
     auto linalgSelect =
-        linalg::SelectOp::create(rewriter, op->getLoc(), ins, outs);
+        linalg::ElementwiseOp::create(rewriter, op->getLoc(), ins, outs,
+                                      mlir::linalg::ElementwiseKind::select);
     linalgOp = linalgSelect;
   }
 
diff --git a/mlir/lib/Dialect/Linalg/Transforms/CMakeLists.txt b/mlir/lib/Dialect/Linalg/Transforms/CMakeLists.txt
index 6dcc6e6d18429..53a8d7b922790 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/CMakeLists.txt
+++ b/mlir/lib/Dialect/Linalg/Transforms/CMakeLists.txt
@@ -10,7 +10,6 @@ add_mlir_dialect_library(MLIRLinalgTransforms
   DropUnitDims.cpp
   ElementwiseOpFusion.cpp
   ElementwiseToLinalg.cpp
-  CategoryToNamedOp.cpp
   EliminateEmptyTensors.cpp
   EraseUnusedOperandsAndResults.cpp
   FoldAddIntoDest.cpp
@@ -27,7 +26,6 @@ add_mlir_dialect_library(MLIRLinalgTransforms
   TransposeMatmul.cpp
   ShardingInterfaceImpl.cpp
   SimplifyDepthwiseConv.cpp
-  NamedToElementwise.cpp
   BlockPackMatmul.cpp
   PackAndUnpackPatterns.cpp
   Padding.cpp
diff --git a/mlir/lib/Dialect/Linalg/Transforms/CategoryToNamedOp.cpp b/mlir/lib/Dialect/Linalg/Transforms/CategoryToNamedOp.cpp
deleted file mode 100644
index 4b36808b8a458..0000000000000
--- a/mlir/lib/Dialect/Linalg/Transforms/CategoryToNamedOp.cpp
+++ /dev/null
@@ -1,61 +0,0 @@
-//===- CategoryToNamedOp.cpp - convert category ops to linalg named ops ---===//
-//
-// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
-// See https://llvm.org/LICENSE.txt for license information.
-// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
-//
-//===----------------------------------------------------------------------===//
-//
-// This file implements rewriting of linalg category ops (e.g.
-// `linalg.elementwise`) to their equivalent named ops (e.g. `linalg.add`).
-// This is the reverse of NamedToElementwise.cpp.
-//
-//===----------------------------------------------------------------------===//
-
-#include "mlir/Dialect/Linalg/IR/Linalg.h"
-#include "mlir/Dialect/Linalg/Transforms/Transforms.h"
-#include "mlir/IR/PatternMatch.h"
-
-using namespace mlir;
-using namespace mlir::linalg;
-
-#define DEBUG_TYPE "linalg-category-to-named"
-
-namespace {
-struct ElementwiseToNamedPattern : public OpRewritePattern<ElementwiseOp> {
-  using OpRewritePattern<ElementwiseOp>::OpRewritePattern;
-
-  LogicalResult matchAndRewrite(ElementwiseOp op,
-                                PatternRewriter &rewriter) const override {
-    // Named elementwise ops only support identity indexing maps.
-    if (!op.getIndexingMapsArray().empty() &&
-        !llvm::all_of(op.getIndexingMapsArray(),
-                      [](AffineMap map) { return map.isIdentity(); }))
-      return failure();
-
-    auto inputs = op.getDpsInputs();
-    auto inits = op.getDpsInits();
-    auto loc = op.getLoc();
-
-    // Helper to create a named op and replace the elementwise op.
-    auto replaceWith = [&](auto namedOp) {
-      using OpTy = decltype(namedOp);
-      rewriter.replaceOp(op, OpTy::create(rewriter, loc, inputs, inits,
-                                          ArrayRef<NamedAttribute>{}));
-      return success();
-    };
-
-    switch (op.getKind()) {
-    case ElementwiseKind::select:
-      return replaceWith(SelectOp{});
-    default:
-      return failure();
-    }
-  }
-};
-} // namespace
-
-void mlir::linalg::populateLinalgCategoryToNamedPatterns(
-    RewritePatternSet &patterns) {
-  patterns.add<ElementwiseToNamedPattern>(patterns.getContext());
-}
diff --git a/mlir/lib/Dialect/Linalg/Transforms/MorphOps.cpp b/mlir/lib/Dialect/Linalg/Transforms/MorphOps.cpp
index 7d360ee734249..02cadf7567de3 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/MorphOps.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/MorphOps.cpp
@@ -44,14 +44,10 @@ void LinalgMorphOpsPass::runOnOperation() {
   RewritePatternSet patterns(&getContext());
 
   // Lowering paths (named -> category -> generic)
-  if (namedToCategory)
-    populateLinalgNamedToElementwisePatterns(patterns);
   if (namedToGeneric || categoryToGeneric)
     populateLinalgNamedOpsGeneralizationPatterns(patterns);
 
   // Lifting paths (named <- category <- generic)
-  if (categoryToNamed)
-    populateLinalgCategoryToNamedPatterns(patterns);
   if (genericToNamed || genericToCategory) {
     GenericOpSpecializationOptions opts;
     opts.emitCategoryOps = genericToCategory;
diff --git a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp b/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
deleted file mode 100644
index d4411d0f7d80a..0000000000000
--- a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
+++ /dev/null
@@ -1,57 +0,0 @@
-//===- NamedToElementwise.cpp - convert linalg named op into elementwise --===//
-//
-// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
-// See https://llvm.org/LICENSE.txt for license information.
-// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
-//
-//===----------------------------------------------------------------------===//
-//
-// This file implements rewriting those linalg named ops that are essentially
-// elementwise e.g. `linalg.add`, to `linalg.elementwise`. This allows further
-// optimization on `linalg.elementwise` such as folding transpose, broadcast.
-//
-//===----------------------------------------------------------------------===//
-
-#include "mlir/Dialect/Linalg/IR/Linalg.h"
-#include "mlir/Dialect/Linalg/Passes.h"
-#include "mlir/Dialect/Linalg/Transforms/Transforms.h"
-#include "mlir/IR/PatternMatch.h"
-#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
-#include "llvm/ADT/SmallVector.h"
-#include "llvm/ADT/TypeSwitch.h"
-
-using namespace mlir;
-using namespace mlir::linalg;
-
-#define DEBUG_TYPE "linalg-named-to-elementwise"
-
-namespace {
-ElementwiseKind getKind(Operation *op) {
-  return llvm::TypeSwitch<Operation *, ElementwiseKind>(op)
-      .Case([](SelectOp) { return ElementwiseKind::select; })
-      .DefaultUnreachable("unhandled case in named to elementwise");
-}
-
-template <typename NamedOpTy>
-struct NamedToElementwisePattern : public OpRewritePattern<NamedOpTy> {
-  using OpRewritePattern<NamedOpTy>::OpRewritePattern;
-
-  LogicalResult matchAndRewrite(NamedOpTy op,
-                                PatternRewriter &rewriter) const override {
-    SmallVector<NamedAttribute> attrs;
-    auto kindAttr = ElementwiseKindAttr::get(op.getContext(), getKind(op));
-    attrs.push_back(rewriter.getNamedAttr("kind", kindAttr));
-    attrs.push_back(
-        rewriter.getNamedAttr("indexing_maps", op.getIndexingMaps()));
-
-    rewriter.replaceOpWithNewOp<ElementwiseOp>(op, op.getDpsInputs(),
-                                               op.getDpsInits(), attrs);
-    return success();
-  }
-};
-} // namespace
-
-void mlir::linalg::populateLinalgNamedToElementwisePatterns(
-    RewritePatternSet &patterns) {
-  patterns.add<NamedToElementwisePattern<SelectOp>>(patterns.getContext());
-}
diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 0c8e5cd212406..88d97d500c10a 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -182,8 +182,7 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
                             scalarBroadcastMap);
       }
       newOp = ElementwiseOp::create(
-          rewriter, genericOp.getLoc(), inputs, genericOp.getDpsInits(),
-          ElementwiseKindAttr::get(rewriter.getContext(), kind),
+          rewriter, genericOp.getLoc(), inputs, genericOp.getDpsInits(), kind,
           rewriter.getAffineMapArrayAttr(indexingMaps));
     }
 
diff --git a/mlir/python/mlir/dialects/linalg/opdsl/ops/core_named_ops.py b/mlir/python/mlir/dialects/linalg/opdsl/ops/core_named_ops.py
index f64526ee558be..16d76bb07dd88 100644
--- a/mlir/python/mlir/dialects/linalg/opdsl/ops/core_named_ops.py
+++ b/mlir/python/mlir/dialects/linalg/opdsl/ops/core_named_ops.py
@@ -21,26 +21,6 @@ def copy(
     O[None] = cast(U, I[None])
 
 
- at linalg_structured_op
-def select(
-    cond=TensorDef(U),
-    lhs=TensorDef(T1),
-    rhs=TensorDef(T1),
-    O=TensorDef(T1, output=True),
-):
-    """Chooses one value based on a binary condition supplied as its first operand.
-
-    The shapes and element types must be identical. The appropriate casts,
-    broadcasts and reductions should be done previously to calling this op.
-
-    This means reduction/broadcast/element cast semantics is explicit. Further
-    passes can take that into account when lowering this code. For example,
-    a `linalg.broadcast` + `linalg.select` sequence can be lowered to a
-    `linalg.generic` with different affine maps for the two operands.
-    """
-    O[None] = TernaryFn.select(cond[None], lhs[None], rhs[None])
-
-
 @linalg_structured_op
 def quantized_matmul(
     A=TensorDef(T1, S.M, S.K),
diff --git a/mlir/test/Dialect/Linalg/elementwise/named-to-elementwise.mlir b/mlir/test/Dialect/Linalg/elementwise/named-to-elementwise.mlir
deleted file mode 100644
index 3da231e729024..0000000000000
--- a/mlir/test/Dialect/Linalg/elementwise/named-to-elementwise.mlir
+++ /dev/null
@@ -1,17 +0,0 @@
-// RUN: mlir-opt %s -linalg-morph-ops=named-to-category -split-input-file | FileCheck %s
-
-// CHECK: @ternary_select(%[[A:.+]]: tensor<4x8x16xi1>, %[[B:.+]]: tensor<4x8x16xf32>, %[[C:.+]]: tensor<4x8x16xf32>)
-// CHECK:   %[[E:.+]] =  tensor.empty() : tensor<4x8x16xf32>
-// CHECK: {{.*}} = linalg.elementwise
-// CHECK-SAME:       kind=#linalg.elementwise_kind<select>
-// CHECK-SAME:       ins(%[[A]], %[[B]], %[[C]] : tensor<4x8x16xi1>, tensor<4x8x16xf32>, tensor<4x8x16xf32>)
-// CHECK-SAME:       outs(%[[E]] : tensor<4x8x16xf32>) -> tensor<4x8x16xf32>
-//
-func.func @ternary_select(%A: tensor<4x8x16xi1>, %B: tensor<4x8x16xf32>, %C: tensor<4x8x16xf32>)
-             -> tensor<4x8x16xf32> {
-  %empty = tensor.empty() : tensor<4x8x16xf32>
-  %select = linalg.select
-              ins(%A, %B, %C : tensor<4x8x16xi1>, tensor<4x8x16xf32>, tensor<4x8x16xf32>)
-              outs(%empty: tensor<4x8x16xf32>) -> tensor<4x8x16xf32>
-  return %select : tensor<4x8x16xf32>
-}
diff --git a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
index 34460d7421d1e..d5a80032f5b8a 100644
--- a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
+++ b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
@@ -404,32 +404,6 @@ func.func @generalize_linalg_map(%arg0: memref<1x8x8x8xf32>, %arg1: memref<1x8x8
 // CHECK:         %[[ADD:.+]] = arith.addf %[[BBARG0]], %[[BBARG1]] : f32
 // CHECK:         linalg.yield %[[ADD]] : f32
 
-// -----
-
-func.func @generalize_select(%cond: memref<7x14x21xi1>, %lhs: memref<7x14x21xf32>, %rhs: memref<7x14x21xf32>,
-                              %out: memref<7x14x21xf32>) {
-  linalg.select ins(%cond, %lhs, %rhs: memref<7x14x21xi1>, memref<7x14x21xf32>, memref<7x14x21xf32>)
-                outs(%out: memref<7x14x21xf32>)
-  return
-}
-
-// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-// CHECK: func @generalize_select
-// CHECK-SAME: (%[[COND:.+]]: memref<7x14x21xi1>, %[[LHS:.+]]: memref<7x14x21xf32>, %[[RHS:.+]]: memref<7x14x21xf32>,
-// CHECK-SAME:  %[[OUT:.+]]: memref<7x14x21xf32>)
-
-// CHECK: linalg.generic
-// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]], #[[MAP]], #[[MAP]]]
-// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
-// CHECK-SAME:  ins(%[[COND]], %[[LHS]], %[[RHS]] : memref<7x14x21xi1>, memref<7x14x21xf32>, memref<7x14x21xf32>)
-// CHECK-SAME: outs(%[[OUT]] : memref<7x14x21xf32>)
-
-// CHECK:         ^{{.+}}(%[[BBARG0:.+]]: i1, %[[BBARG1:.+]]: f32, %[[BBARG2:.+]]: f32, %[[BBARG3:.+]]: f32)
-// CHECK-NEXT:      %[[select:.+]] = arith.select %[[BBARG0]], %[[BBARG1]], %[[BBARG2]] : f32
-// CHECK-NEXT:      linalg.yield %[[select]] : f32
-
-
 // -----
 
 // CHECK-LABEL: func @fill_tensor
diff --git a/mlir/test/Dialect/Linalg/linalg-morph-elementwise-to-named.mlir b/mlir/test/Dialect/Linalg/linalg-morph-elementwise-to-named.mlir
index 16206c7e837c8..2baae9a9714c7 100644
--- a/mlir/test/Dialect/Linalg/linalg-morph-elementwise-to-named.mlir
+++ b/mlir/test/Dialect/Linalg/linalg-morph-elementwise-to-named.mlir
@@ -112,24 +112,6 @@ func.func @binary_ops_float(%A: tensor<?x?xf32>, %B: tensor<?x?xf32>,
 
 // -----
 
-func.func @ternary_select(%A: tensor<?x?xi1>, %B: tensor<?x?xf32>,
-                          %C: tensor<?x?xf32>,
-                          %Out: tensor<?x?xf32>) -> tensor<?x?xf32> {
-  %0 = linalg.elementwise kind=#linalg.elementwise_kind<select>
-    ins(%A, %B, %C : tensor<?x?xi1>, tensor<?x?xf32>, tensor<?x?xf32>)
-    outs(%Out : tensor<?x?xf32>) -> tensor<?x?xf32>
-  return %0 : tensor<?x?xf32>
-}
-
-// CHECK-LABEL: ternary_select
-// CHECK-SAME: %[[A:.+]]: tensor<?x?xi1>, %[[B:.+]]: tensor<?x?xf32>, %[[C:.+]]: tensor<?x?xf32>, %[[OUT:.+]]: tensor<?x?xf32>)
-// CHECK-NOT: linalg.elementwise
-// CHECK: linalg.select
-// CHECK-SAME: ins(%[[A]], %[[B]], %[[C]] : tensor<?x?xi1>, tensor<?x?xf32>, tensor<?x?xf32>)
-// CHECK-SAME: outs(%[[OUT]] : tensor<?x?xf32>) -> tensor<?x?xf32>
-
-// -----
-
 // Non-identity indexing maps: should NOT be converted to named op.
 func.func @non_identity_maps(%A: tensor<?xf32>, %Out: tensor<?x?xf32>) -> tensor<?x?xf32> {
   %0 = linalg.elementwise
diff --git a/mlir/test/Dialect/Linalg/named-ops-fail.mlir b/mlir/test/Dialect/Linalg/named-ops-fail.mlir
deleted file mode 100644
index b1809e652aabe..0000000000000
--- a/mlir/test/Dialect/Linalg/named-ops-fail.mlir
+++ /dev/null
@@ -1,48 +0,0 @@
-// RUN: not mlir-opt -split-input-file -verify-diagnostics %s 2>&1 | FileCheck %s
-
-func.func @select_type_cast(%arg0: memref<4x8x16xi1>, %arg1: memref<4x8x16xf16>, %arg2: memref<4x8x16xf32>, %arg3: memref<4x8x16xf32>) {
-  // CHECK: op failed to verify that all of {true_value, false_value, result} have same type
-  linalg.select ins(%arg0, %arg1, %arg2 : memref<4x8x16xi1>, memref<4x8x16xf16>, memref<4x8x16xf32>) outs(%arg3: memref<4x8x16xf32>)
-  return
-}
-
-// -----
-
-func.func @select_wrong_condition_type(%arg0: memref<4x8x16xf32>, %arg1: memref<4x8x16xf32>, %arg2: memref<4x8x16xf32>, %arg3: memref<4x8x16xf32>) {
-  // CHECK: op operand #0 must be bool-like, but got 'f32'
-  linalg.select ins(%arg0, %arg1, %arg2 : memref<4x8x16xf32>, memref<4x8x16xf32>, memref<4x8x16xf32>) outs(%arg3: memref<4x8x16xf32>)
-  return
-}
-
-// -----
-
-// linalg.select with all-integer operands
-func.func @select_all_integer(%arg0: memref<4x8x16xi32>, %arg1: memref<4x8x16xi32>, %arg2: memref<4x8x16xi32>, %arg3: memref<4x8x16xi32>) {
-  // CHECK: op operand #0 must be bool-like, but got 'i32'
-  linalg.select ins(%arg0, %arg1, %arg2 : memref<4x8x16xi32>, memref<4x8x16xi32>, memref<4x8x16xi32>) outs(%arg3: memref<4x8x16xi32>)
-  return
-}
-
-// -----
-
-// Regression test: linalg.select with index type operands should emit a
-// diagnostic instead of crashing (https://github.com/llvm/llvm-project/issues/179046).
-func.func @select_invalid_index_type(%cond: index, %a: index, %b: index,
-                                     %out: tensor<1xindex>) -> tensor<1xindex> {
-  // CHECK: op operand #0 must be bool-like, but got 'index'
-  %0 = linalg.select ins(%cond, %a, %b : index, index, index)
-                     outs(%out : tensor<1xindex>) -> tensor<1xindex>
-  return %0 : tensor<1xindex>
-}
-
-// -----
-
-// linalg.select with an integer (non-i1) condition and floating-point values:
-func.func @select_invalid_integer_cond_float_values(%cond: tensor<4xi32>,
-    %a: tensor<4xf32>, %b: tensor<4xf32>,
-    %out: tensor<4xf32>) -> tensor<4xf32> {
-// CHECK: op operand #0 must be bool-like, but got 'i32'
-  %0 = linalg.select ins(%cond, %a, %b : tensor<4xi32>, tensor<4xf32>, tensor<4xf32>)
-                     outs(%out : tensor<4xf32>) -> tensor<4xf32>
-  return %0 : tensor<4xf32>
-}
diff --git a/mlir/test/Dialect/Linalg/named-ops.mlir b/mlir/test/Dialect/Linalg/named-ops.mlir
index f0003f93bc2e0..144000f412172 100644
--- a/mlir/test/Dialect/Linalg/named-ops.mlir
+++ b/mlir/test/Dialect/Linalg/named-ops.mlir
@@ -2009,54 +2009,6 @@ func.func @fill_tensor(%f: f32, %v: vector<2x4xf32>) -> (tensor<f32>, tensor<vec
   return %0, %1: tensor<f32>, tensor<vector<2x4xf32>>
 }
 
-// -----
-
-// CHECK-LABEL: func @select_dynamic
-func.func @select_dynamic(%arg0: memref<?x?x?xi1>, %arg1: memref<?x?x?xf32>, %arg2: memref<?x?x?xf32>, %arg3: memref<?x?x?xf32>) {
-  // CHECK: linalg.select
-  // CHECK-SAME: ins(%{{.+}}, %{{.+}}, %{{.+}} : memref<?x?x?xi1>, memref<?x?x?xf32>, memref<?x?x?xf32>)
-  // CHECK-SAME: outs(%{{.+}} : memref<?x?x?xf32>)
-  linalg.select ins(%arg0, %arg1, %arg2 : memref<?x?x?xi1>, memref<?x?x?xf32>, memref<?x?x?xf32>) outs(%arg3: memref<?x?x?xf32>)
-  return
-}
-
-// -----
-
-// CHECK-LABEL: func @select_static
-func.func @select_static(%arg0: memref<4x8x16xi1>, %arg1: memref<4x8x16xf32>, %arg2: memref<4x8x16xf32>, %arg3: memref<4x8x16xf32>) {
-  // CHECK: linalg.select
-  // CHECK-SAME: ins(%{{.+}}, %{{.+}}, %{{.+}} : memref<4x8x16xi1>, memref<4x8x16xf32>, memref<4x8x16xf32>)
-  // CHECK-SAME: outs(%{{.+}} : memref<4x8x16xf32>)
-  linalg.select ins(%arg0, %arg1, %arg2 : memref<4x8x16xi1>, memref<4x8x16xf32>, memref<4x8x16xf32>) outs(%arg3: memref<4x8x16xf32>)
-  return
-}
-
-// -----
-
-// CHECK-LABEL: func @select_tensor
-func.func @select_tensor(%arg0: tensor<4x8x16xi1>, %arg1: tensor<4x8x16xf32>, %arg2: tensor<4x8x16xf32>) -> tensor<4x8x16xf32> {
-  %0 = tensor.empty() : tensor<4x8x16xf32>
-  // CHECK: linalg.select
-  // CHECK-SAME: ins(%{{.+}}, %{{.+}}, %{{.+}} : tensor<4x8x16xi1>, tensor<4x8x16xf32>, tensor<4x8x16xf32>)
-  // CHECK-SAME: outs(%{{.+}} : tensor<4x8x16xf32>)
-  %1 = linalg.select ins(%arg0, %arg1, %arg2 : tensor<4x8x16xi1>, tensor<4x8x16xf32>, tensor<4x8x16xf32>) outs(%0: tensor<4x8x16xf32>) -> tensor<4x8x16xf32>
-  return %1 : tensor<4x8x16xf32>
-}
-
-// -----
-
-// CHECK-LABEL: func @select_integer_values
-// linalg.select with i1 condition and integer values: headBool=true (i1 bitwidth==1)
-// → valid, arith.select accepts i1 as condition regardless of value types.
-func.func @select_integer_values(%arg0: tensor<4x8x16xi1>, %arg1: tensor<4x8x16xi32>, %arg2: tensor<4x8x16xi32>) -> tensor<4x8x16xi32> {
-  %0 = tensor.empty() : tensor<4x8x16xi32>
-  // CHECK: linalg.select
-  // CHECK-SAME: ins(%{{.+}}, %{{.+}}, %{{.+}} : tensor<4x8x16xi1>, tensor<4x8x16xi32>, tensor<4x8x16xi32>)
-  // CHECK-SAME: outs(%{{.+}} : tensor<4x8x16xi32>)
-  %1 = linalg.select ins(%arg0, %arg1, %arg2 : tensor<4x8x16xi1>, tensor<4x8x16xi32>, tensor<4x8x16xi32>) outs(%0: tensor<4x8x16xi32>) -> tensor<4x8x16xi32>
-  return %1 : tensor<4x8x16xi32>
-}
-
 //===----------------------------------------------------------------------===//
 // linalg.pack + linalg.unpack
 //===----------------------------------------------------------------------===//

>From 55033334cf929a7de6e8352583f588c585e32449 Mon Sep 17 00:00:00 2001
From: rengolin <rengolin at systemcall.eu>
Date: Thu, 3 Sep 2026 07:20:56 -0700
Subject: [PATCH 2/3] [MLIR][Linalg] Named ops morph cleanup

---
 .../Dialect/Linalg/Transforms/Specialize.cpp  | 196 +++++++++---------
 ...ic-by-unfolding-projected-permutation.mlir |   2 +-
 .../Linalg/specialize-generic-ops.mlir        |  66 +-----
 ...ansform-op-specialize-elemwise-binary.mlir |   2 +-
 4 files changed, 108 insertions(+), 158 deletions(-)

diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 88d97d500c10a..8bec5e1c8b137 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -158,7 +158,7 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
     using NamedOpTy = decltype(namedOp);
     // A null named op means the op only has a category form; emit
     // `linalg.elementwise` regardless of the requested output form.
-    if (!emitCategoryOp && !std::is_null_pointer_v<NamedOpTy>) {
+    if (!std::is_null_pointer_v<NamedOpTy>) {
       if constexpr (!std::is_null_pointer_v<NamedOpTy>)
         newOp = NamedOpTy::create(rewriter, genericOp.getLoc(), inputs,
                                   genericOp.getDpsInits(),
@@ -190,108 +190,106 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
     return newOp;
   };
 
-  if (isUnary) {
-    // Unary ops only have the category (elementwise) form; the linalg.* named
-    // unary ops have been removed.
-    if (emitCategoryOp) {
-      if (isa<math::ExpOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::exp);
-      if (isa<math::AbsFOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::abs);
-      if (isa<math::CeilOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::ceil);
-      if (isa<math::FloorOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::floor);
-      if (isa<arith::NegFOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::negf);
-      if (auto divOp = dyn_cast<arith::DivFOp>(op)) {
-        if (auto constOp = dyn_cast_if_present<arith::ConstantOp>(
-                divOp.getLhs().getDefiningOp()))
-          if (cast<FloatAttr>(constOp.getValue())
-                  .getValue()
-                  .isExactlyValue(1.0))
-            return replaceOp(nullptr, ElementwiseKind::reciprocal,
-                             /*mayHoistScalarOperand=*/false);
-      }
-      if (isa<math::RoundOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::round);
-      if (isa<math::SqrtOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::sqrt);
-      if (isa<math::RsqrtOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::rsqrt);
-      if (auto mulOp = dyn_cast<arith::MulFOp>(op);
-          mulOp && mulOp.getLhs() == mulOp.getRhs())
-        return replaceOp(nullptr, ElementwiseKind::square);
-      if (isa<math::TanhOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::tanh);
-      if (isa<math::ErfOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::erf);
-      if (isa<math::SinOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::sin);
-      if (isa<math::CosOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::cos);
-      if (isa<math::TanOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::tan);
-      if (isa<math::AcosOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::acos);
-      if (isa<math::AcoshOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::acosh);
-      if (isa<math::AsinOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::asin);
-      if (isa<math::AsinhOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::asinh);
-      if (isa<math::AtanOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::atan);
-      if (isa<math::AtanhOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::atanh);
-      if (isa<math::LogOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::log);
-      if (isa<math::Log10Op>(op))
-        return replaceOp(nullptr, ElementwiseKind::log10);
-      if (isa<math::Log1pOp>(op))
-        return replaceOp(nullptr, ElementwiseKind::log1p);
-      if (isa<math::Log2Op>(op))
-        return replaceOp(nullptr, ElementwiseKind::log2);
+  // There are no named ops for these elementwise operations; can only emit the
+  // category form.
+  if (emitCategoryOp) {
+    if (isa<math::ExpOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::exp);
+    if (isa<math::AbsFOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::abs);
+    if (isa<math::CeilOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::ceil);
+    if (isa<math::FloorOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::floor);
+    if (isa<arith::NegFOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::negf);
+    if (auto divOp = dyn_cast<arith::DivFOp>(op)) {
+      if (auto constOp = dyn_cast_if_present<arith::ConstantOp>(
+              divOp.getLhs().getDefiningOp()))
+        if (cast<FloatAttr>(constOp.getValue())
+                .getValue()
+                .isExactlyValue(1.0))
+          return replaceOp(nullptr, ElementwiseKind::reciprocal,
+                            /*mayHoistScalarOperand=*/false);
     }
-
-    // At this point, we exhaustively checked the available unary named ops. The
-    // 1-input generic op might be representable as a `linalg.elementwise` that
-    // broadcasts a scalar operand. But if we can't emit the category op or
-    // don't have a scalar operand, exit now.
-    if (!emitCategoryOp || !hasScalarOperand)
+    if (isa<math::RoundOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::round);
+    if (isa<math::SqrtOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::sqrt);
+    if (isa<math::RsqrtOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::rsqrt);
+    if (auto mulOp = dyn_cast<arith::MulFOp>(op);
+        mulOp && mulOp.getLhs() == mulOp.getRhs())
+      return replaceOp(nullptr, ElementwiseKind::square);
+    if (isa<math::TanhOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::tanh);
+    if (isa<math::ErfOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::erf);
+    if (isa<math::SinOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::sin);
+    if (isa<math::CosOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::cos);
+    if (isa<math::TanOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::tan);
+    if (isa<math::AcosOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::acos);
+    if (isa<math::AcoshOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::acosh);
+    if (isa<math::AsinOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::asin);
+    if (isa<math::AsinhOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::asinh);
+    if (isa<math::AtanOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::atan);
+    if (isa<math::AtanhOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::atanh);
+    if (isa<math::LogOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::log);
+    if (isa<math::Log10Op>(op))
+      return replaceOp(nullptr, ElementwiseKind::log10);
+    if (isa<math::Log1pOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::log1p);
+    if (isa<math::Log2Op>(op))
+      return replaceOp(nullptr, ElementwiseKind::log2);
+
+    // The remaining kinds are binary. A single-input generic can only be
+    // represented as a binary elementwise if it has a scalar operand to hoist;
+    // otherwise (e.g. a body reusing a block argument twice) it has no
+    // category form.
+    if (isUnary && !hasScalarOperand)
       return rewriter.notifyMatchFailure(
-          genericOp, "unary elementwise operation cannot be specialized to "
-                     "named or category op");
+          genericOp, "unary elementwise operation cannot be specialized to a "
+                     "category op");
+
+    // Boolean-typed `linalg.add` and `linalg.mul` require special handling.
+    bool allBool = llvm::all_of(op->getOperands(),
+                                [](Value v) { return v.getType().isInteger(1); });
+
+    if (isa<arith::AddFOp, arith::AddIOp, complex::AddOp>(op) ||
+        (allBool && isa<arith::OrIOp>(op)))
+      return replaceOp(nullptr, ElementwiseKind::add);
+    if (isa<arith::SubIOp, arith::SubFOp, complex::SubOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::sub);
+    if (isa<arith::MulIOp, arith::MulFOp, complex::MulOp>(op) ||
+        (allBool && isa<arith::AndIOp>(op)))
+      return replaceOp(nullptr, ElementwiseKind::mul);
+    if (isa<arith::DivSIOp, arith::DivFOp, complex::DivOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::div);
+    if (isa<arith::DivUIOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::div_unsigned);
+    if (isa<arith::MaxSIOp, arith::MaximumFOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::max_signed);
+    if (isa<arith::MinSIOp, arith::MinimumFOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::min_signed);
+    if (isa<math::PowFOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::powf);
+    // No named ops for unsigned maximum/minimum.
+    if (isa<arith::MaxUIOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::max_unsigned);
+    if (isa<arith::MinUIOp>(op))
+      return replaceOp(nullptr, ElementwiseKind::min_unsigned);
   }
 
-  // Boolean-typed `linalg.add` and `linalg.mul` require special handling.
-  bool allBool = llvm::all_of(op->getOperands(),
-                              [](Value v) { return v.getType().isInteger(1); });
-
-  if (isa<arith::AddFOp, arith::AddIOp, complex::AddOp>(op) ||
-      (allBool && isa<arith::OrIOp>(op)))
-    return replaceOp(nullptr, ElementwiseKind::add);
-  if (isa<arith::SubIOp, arith::SubFOp, complex::SubOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::sub);
-  if (isa<arith::MulIOp, arith::MulFOp, complex::MulOp>(op) ||
-      (allBool && isa<arith::AndIOp>(op)))
-    return replaceOp(nullptr, ElementwiseKind::mul);
-  if (isa<arith::DivSIOp, arith::DivFOp, complex::DivOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::div);
-  if (isa<arith::DivUIOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::div_unsigned);
-  if (isa<arith::MaxSIOp, arith::MaximumFOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::max_signed);
-  if (isa<arith::MinSIOp, arith::MinimumFOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::min_signed);
-  if (isa<math::PowFOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::powf);
-  // No named ops for unsigned maximum/minimum.
-  if (isa<arith::MaxUIOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::max_unsigned);
-  if (isa<arith::MinUIOp>(op))
-    return replaceOp(nullptr, ElementwiseKind::min_unsigned);
-
   return rewriter.notifyMatchFailure(
       genericOp,
       "elementwise operation cannot be specialized to named or category op");
diff --git a/mlir/test/Dialect/Linalg/decompose-generic-by-unfolding-projected-permutation.mlir b/mlir/test/Dialect/Linalg/decompose-generic-by-unfolding-projected-permutation.mlir
index f9c986986eb33..3cc39c5ac8333 100644
--- a/mlir/test/Dialect/Linalg/decompose-generic-by-unfolding-projected-permutation.mlir
+++ b/mlir/test/Dialect/Linalg/decompose-generic-by-unfolding-projected-permutation.mlir
@@ -1,4 +1,4 @@
-// RUN: mlir-opt %s -split-input-file --linalg-specialize-generic-ops | FileCheck %s
+// RUN: mlir-opt %s -split-input-file --linalg-specialize-generic-ops -linalg-morph-ops=generic-to-category | FileCheck %s
 
 #projection = affine_map<(d0, d1, d2, d3, d4) -> (d2, d3, d1)>
 #identity   = affine_map<(d0, d1, d2, d3, d4) -> (d0, d1, d2, d3, d4)>
diff --git a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
index ddd3f4a543f85..2ae01dff10107 100644
--- a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
+++ b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
@@ -324,25 +324,8 @@ func.func @binary_ops_int(%A: tensor<?x?xi32>, %B: tensor<?x?xi32>,
 // ALL-SAME: %[[A:.+]]: [[TTY:tensor<\?x\?xi32>]], %[[B:.+]]: [[TTY]],
 // ALL-SAME: %[[OUT:.+]]: [[TTY]]) -> [[TTY]]
 
-// NAMED-NOT: linalg.generic
-// NAMED: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
-// NAMED-SAME: ins(%[[A]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES2:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
-// NAMED-SAME: ins(%[[RES1]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES3:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<div>
-// NAMED-SAME: ins(%[[RES2]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES4:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<div_unsigned>
-// NAMED-SAME: ins(%[[RES3]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES5:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<max_signed>
-// NAMED-SAME: ins(%[[RES4]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES6:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<min_signed>
-// NAMED-SAME: ins(%[[RES5]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
+// NAMED-NOT: linalg.elementwise
+// NAMED: linalg.generic
 
 // CATEGORY-NOT: linalg.generic
 // CATEGORY: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
@@ -430,25 +413,8 @@ func.func @binary_ops_float(%A: tensor<?x?xf32>, %B: tensor<?x?xf32>,
 // ALL-SAME: %[[A:.+]]: [[TTY:tensor<\?x\?xf32>]], %[[B:.+]]: [[TTY]],
 // ALL-SAME: %[[OUT:.+]]: [[TTY]]) -> [[TTY]]
 
-// NAMED-NOT: linalg.generic
-// NAMED: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
-// NAMED-SAME: ins(%[[A]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES2:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
-// NAMED-SAME: ins(%[[RES1]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES3:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<div>
-// NAMED-SAME: ins(%[[RES2]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES4:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<max_signed>
-// NAMED-SAME: ins(%[[RES3]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES5:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<min_signed>
-// NAMED-SAME: ins(%[[RES4]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES6:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<powf>
-// NAMED-SAME: ins(%[[RES5]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
+// NAMED-NOT: linalg.elementwise
+// NAMED: linalg.generic
 
 // CATEGORY-NOT: linalg.generic
 // CATEGORY: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
@@ -511,16 +477,8 @@ func.func @binary_ops_complex(%A: tensor<?x?xcomplex<f32>>,
 // ALL-SAME: %[[A:.+]]: [[TTY:tensor<\?x\?xcomplex<f32>>]], %[[B:.+]]: [[TTY]],
 // ALL-SAME: %[[OUT:.+]]: [[TTY]]) -> [[TTY]]
 
-// NAMED-NOT: linalg.generic
-// NAMED: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
-// NAMED-SAME: ins(%[[A]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES2:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
-// NAMED-SAME: ins(%[[RES1]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED: %[[RES3:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<div>
-// NAMED-SAME: ins(%[[RES2]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
+// NAMED-NOT: linalg.elementwise
+// NAMED: linalg.generic
 
 // CATEGORY-NOT: linalg.generic
 // CATEGORY: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sub>
@@ -554,10 +512,8 @@ func.func @binary_ops_bool(%A: tensor<?x?xi1>, %B: tensor<?x?xi1>,
 // ALL-SAME: %[[A:.+]]: [[TTY:tensor<\?x\?xi1>]], %[[B:.+]]: [[TTY]],
 // ALL-SAME: %[[OUT:.+]]: [[TTY]]) -> [[TTY]]
 
-// NAMED-NOT: linalg.generic
-// NAMED: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
-// NAMED-SAME: ins(%[[A]], %[[B]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
+// NAMED-NOT: linalg.elementwise
+// NAMED: linalg.generic
 
 // CATEGORY-NOT: linalg.generic
 // CATEGORY: %[[RES1:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
@@ -663,12 +619,8 @@ func.func @binary_ops_swapped(%A: tensor<?x?xf32>, %B: tensor<?x?xf32>,
 // ALL-SAME: %[[C:.+]]: [[TTY1D:tensor<\?xf32>]],
 // ALL-SAME: %[[OUT:.+]]: [[TTY]]) -> [[TTY]]
 
-// NAMED: %[[RES0:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
-// NAMED-SAME: ins(%[[B]], %[[A]] : [[TTY]], [[TTY]])
-// NAMED-SAME: outs(%[[OUT]] : [[TTY]]) -> [[TTY]]
-// NAMED-NOT: linalg.sub
+// NAMED-NOT: linalg.elementwise
 // NAMED: linalg.generic
-// NAMED-SAME: ins(%[[RES0]], %[[C]] : [[TTY]], [[TTY1D]])
 
 // CATEGORY-NOT: linalg.generic
 // CATEGORY: %[[RES0:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<mul>
diff --git a/mlir/test/Dialect/Linalg/transform-op-specialize-elemwise-binary.mlir b/mlir/test/Dialect/Linalg/transform-op-specialize-elemwise-binary.mlir
index 943ebd9ba0be4..dbab2dfcf4692 100644
--- a/mlir/test/Dialect/Linalg/transform-op-specialize-elemwise-binary.mlir
+++ b/mlir/test/Dialect/Linalg/transform-op-specialize-elemwise-binary.mlir
@@ -225,7 +225,7 @@ func.func @specialize_sub_swapped_operands(%arg0: tensor<?x?xf32>, %arg1: tensor
 module attributes {transform.with_named_sequence} {
   transform.named_sequence @__transform_main(%arg0: !transform.any_op {transform.readonly}) {
     %0 = transform.structured.match interface{LinalgOp} in %arg0 : (!transform.any_op) -> !transform.any_op
-    %1 = transform.structured.specialize %0 : (!transform.any_op) -> !transform.any_op
+    %1 = transform.structured.specialize %0 {emit_category = true} : (!transform.any_op) -> !transform.any_op
     transform.yield
   }
 }

>From 4cb33d636ce1b18e311fee5e9870f40dccb323e3 Mon Sep 17 00:00:00 2001
From: rengolin <rengolin at systemcall.eu>
Date: Thu, 3 Sep 2026 07:28:27 -0700
Subject: [PATCH 3/3] format

---
 mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp | 10 ++++------
 1 file changed, 4 insertions(+), 6 deletions(-)

diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 8bec5e1c8b137..587aecc7adaf1 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -206,11 +206,9 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
     if (auto divOp = dyn_cast<arith::DivFOp>(op)) {
       if (auto constOp = dyn_cast_if_present<arith::ConstantOp>(
               divOp.getLhs().getDefiningOp()))
-        if (cast<FloatAttr>(constOp.getValue())
-                .getValue()
-                .isExactlyValue(1.0))
+        if (cast<FloatAttr>(constOp.getValue()).getValue().isExactlyValue(1.0))
           return replaceOp(nullptr, ElementwiseKind::reciprocal,
-                            /*mayHoistScalarOperand=*/false);
+                           /*mayHoistScalarOperand=*/false);
     }
     if (isa<math::RoundOp>(op))
       return replaceOp(nullptr, ElementwiseKind::round);
@@ -262,8 +260,8 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
                      "category op");
 
     // Boolean-typed `linalg.add` and `linalg.mul` require special handling.
-    bool allBool = llvm::all_of(op->getOperands(),
-                                [](Value v) { return v.getType().isInteger(1); });
+    bool allBool = llvm::all_of(
+        op->getOperands(), [](Value v) { return v.getType().isInteger(1); });
 
     if (isa<arith::AddFOp, arith::AddIOp, complex::AddOp>(op) ||
         (allBool && isa<arith::OrIOp>(op)))



More information about the llvm-branch-commits mailing list