[Mlir-commits] [mlir] [mlir][linalg] Add sin, cos, tan to elementwise operations (PR #200950)
Vinit Deodhar
llvmlistbot at llvm.org
Mon Jun 8 06:29:44 PDT 2026
https://github.com/vinitdeodhar updated https://github.com/llvm/llvm-project/pull/200950
>From d621199b3897adb41328726e1e8302dda3dd85c7 Mon Sep 17 00:00:00 2001
From: Vinit Deodhar <vinitdeodhar at users.noreply.github.com>
Date: Mon, 1 Jun 2026 17:28:20 -0400
Subject: [PATCH 1/3] [mlir][linalg] Add sin, cos, tan to elementwise
operations
Add sin, cos, and tan as UnaryFn entries in the linalg dialect,
enabling their use via linalg.elementwise, named ops (linalg.sin,
linalg.cos, linalg.tan), and specialization from linalg.generic.
Co-Authored-By: Claude Opus 4.6 <noreply at anthropic.com>
---
.../mlir/Dialect/Linalg/IR/LinalgEnums.td | 5 +-
.../Linalg/IR/LinalgNamedStructuredOps.yaml | 105 ++++++++++++++++++
mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp | 6 +
.../Linalg/Transforms/NamedToElementwise.cpp | 6 +
.../Dialect/Linalg/Transforms/Specialize.cpp | 6 +
.../Dialect/Linalg/generalize-named-ops.mlir | 63 +++++++++++
...oundtrip-morphism-linalg-category-ops.mlir | 15 +++
.../Linalg/specialize-generic-ops.mlir | 41 ++++++-
8 files changed, 245 insertions(+), 2 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgEnums.td b/mlir/include/mlir/Dialect/Linalg/IR/LinalgEnums.td
index 1109db973f522..f6b5edd326c85 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgEnums.td
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgEnums.td
@@ -29,7 +29,10 @@ def UnaryFn : I32EnumAttr<"UnaryFn", "", [
I32EnumAttrCase<"rsqrt", 9>,
I32EnumAttrCase<"square", 10>,
I32EnumAttrCase<"tanh", 11>,
- I32EnumAttrCase<"erf", 12>
+ I32EnumAttrCase<"erf", 12>,
+ I32EnumAttrCase<"sin", 13>,
+ I32EnumAttrCase<"cos", 14>,
+ I32EnumAttrCase<"tan", 15>
]> {
let genSpecializedAttr = 0;
let cppNamespace = "::mlir::linalg";
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
index 521afc991063f..f62ffe113d123 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
@@ -499,6 +499,111 @@ structured_op: !LinalgStructuredOpConfig
- !ScalarExpression
scalar_arg: I
--- !LinalgOpConfig
+metadata: !LinalgOpMetadata
+ name: sin
+ cpp_class_name: SinOp
+ doc: |-
+ Applies sin(x) elementwise.
+
+ No numeric casting is performed on the input operand.
+structured_op: !LinalgStructuredOpConfig
+ args:
+ - !LinalgOperandDefConfig
+ name: I
+ 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<() -> ()>
+ iterator_types: []
+ assignments:
+ - !ScalarAssign
+ arg: O
+ value: !ScalarExpression
+ scalar_fn:
+ kind: unary
+ fn_name: sin
+ operands:
+ - !ScalarExpression
+ scalar_arg: I
+--- !LinalgOpConfig
+metadata: !LinalgOpMetadata
+ name: cos
+ cpp_class_name: CosOp
+ doc: |-
+ Applies cos(x) elementwise.
+
+ No numeric casting is performed on the input operand.
+structured_op: !LinalgStructuredOpConfig
+ args:
+ - !LinalgOperandDefConfig
+ name: I
+ 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<() -> ()>
+ iterator_types: []
+ assignments:
+ - !ScalarAssign
+ arg: O
+ value: !ScalarExpression
+ scalar_fn:
+ kind: unary
+ fn_name: cos
+ operands:
+ - !ScalarExpression
+ scalar_arg: I
+--- !LinalgOpConfig
+metadata: !LinalgOpMetadata
+ name: tan
+ cpp_class_name: TanOp
+ doc: |-
+ Applies tan(x) elementwise.
+
+ No numeric casting is performed on the input operand.
+structured_op: !LinalgStructuredOpConfig
+ args:
+ - !LinalgOperandDefConfig
+ name: I
+ 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<() -> ()>
+ iterator_types: []
+ assignments:
+ - !ScalarAssign
+ arg: O
+ value: !ScalarExpression
+ scalar_fn:
+ kind: unary
+ fn_name: tan
+ operands:
+ - !ScalarExpression
+ scalar_arg: I
+--- !LinalgOpConfig
metadata: !LinalgOpMetadata
name: add
cpp_class_name: AddOp
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
index de7f4d1610bd0..6b61564eb5fcc 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
@@ -503,6 +503,12 @@ class RegionBuilderHelper {
return math::TanhOp::create(builder, arg.getLoc(), arg);
case UnaryFn::erf:
return math::ErfOp::create(builder, arg.getLoc(), arg);
+ case UnaryFn::sin:
+ return math::SinOp::create(builder, arg.getLoc(), arg);
+ case UnaryFn::cos:
+ return math::CosOp::create(builder, arg.getLoc(), arg);
+ case UnaryFn::tan:
+ return math::TanOp::create(builder, arg.getLoc(), arg);
}
if (emitError) {
emitError() << "unsupported unary function";
diff --git a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp b/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
index c9045566473cb..02f6677cfbf7e 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
@@ -48,6 +48,9 @@ ElementwiseKind getKind(Operation *op) {
.Case([](SquareOp) { return ElementwiseKind::square; })
.Case([](TanhOp) { return ElementwiseKind::tanh; })
.Case([](ErfOp) { return ElementwiseKind::erf; })
+ .Case([](SinOp) { return ElementwiseKind::sin; })
+ .Case([](CosOp) { return ElementwiseKind::cos; })
+ .Case([](TanOp) { return ElementwiseKind::tan; })
.DefaultUnreachable("unhandled case in named to elementwise");
}
@@ -92,4 +95,7 @@ void mlir::linalg::populateLinalgNamedToElementwisePatterns(
patterns.add<NamedToElementwisePattern<SquareOp>>(patterns.getContext());
patterns.add<NamedToElementwisePattern<TanhOp>>(patterns.getContext());
patterns.add<NamedToElementwisePattern<ErfOp>>(patterns.getContext());
+ patterns.add<NamedToElementwisePattern<SinOp>>(patterns.getContext());
+ patterns.add<NamedToElementwisePattern<CosOp>>(patterns.getContext());
+ patterns.add<NamedToElementwisePattern<TanOp>>(patterns.getContext());
}
diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 2ad2774ff1a0c..25f69974da29c 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -223,6 +223,12 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
return replaceOp(TanhOp{}, ElementwiseKind::tanh);
if (isa<math::ErfOp>(op))
return replaceOp(ErfOp{}, ElementwiseKind::erf);
+ if (isa<math::SinOp>(op))
+ return replaceOp(SinOp{}, ElementwiseKind::sin);
+ if (isa<math::CosOp>(op))
+ return replaceOp(CosOp{}, ElementwiseKind::cos);
+ if (isa<math::TanOp>(op))
+ return replaceOp(TanOp{}, ElementwiseKind::tan);
// At this point, we exhaustively checked the available unary named ops. The
// 1-input generic op might be representable as a `linalg.elementwise` that
diff --git a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
index e346bee901f1d..21d0d103774a6 100644
--- a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
+++ b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
@@ -806,6 +806,69 @@ func.func @generalize_erf(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>)
// -----
+func.func @generalize_sin(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
+ linalg.sin ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
+ return
+}
+
+// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+// CHECK: func @generalize_sin
+// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
+
+// CHECK: linalg.generic
+// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
+// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
+// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
+
+// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
+// CHECK-NEXT: %[[sin:.+]] = math.sin %[[BBARG0]] : f32
+// CHECK-NEXT: linalg.yield %[[sin]] : f32
+
+// -----
+
+func.func @generalize_cos(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
+ linalg.cos ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
+ return
+}
+
+// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+// CHECK: func @generalize_cos
+// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
+
+// CHECK: linalg.generic
+// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
+// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
+// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
+
+// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
+// CHECK-NEXT: %[[cos:.+]] = math.cos %[[BBARG0]] : f32
+// CHECK-NEXT: linalg.yield %[[cos]] : f32
+
+// -----
+
+func.func @generalize_tan(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
+ linalg.tan ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
+ return
+}
+
+// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
+
+// CHECK: func @generalize_tan
+// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
+
+// CHECK: linalg.generic
+// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
+// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
+// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
+
+// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
+// CHECK-NEXT: %[[tan:.+]] = math.tan %[[BBARG0]] : f32
+// CHECK-NEXT: linalg.yield %[[tan]] : f32
+
+// -----
+
func.func @generalize_max(%lhs: memref<7x14x21xf32>, %rhs: memref<7x14x21xf32>,
%out: memref<7x14x21xf32>) {
linalg.max ins(%lhs, %rhs : memref<7x14x21xf32>, memref<7x14x21xf32>)
diff --git a/mlir/test/Dialect/Linalg/roundtrip-morphism-linalg-category-ops.mlir b/mlir/test/Dialect/Linalg/roundtrip-morphism-linalg-category-ops.mlir
index f63b795108cdb..5302072777b9f 100644
--- a/mlir/test/Dialect/Linalg/roundtrip-morphism-linalg-category-ops.mlir
+++ b/mlir/test/Dialect/Linalg/roundtrip-morphism-linalg-category-ops.mlir
@@ -32,6 +32,12 @@ func.func @unary_ops(%A: memref<7x14x21xf32>, %Out: memref<7x14x21xf32>) {
ins(%A : memref<7x14x21xf32>) outs(%Out : memref<7x14x21xf32>)
linalg.elementwise kind=#linalg.elementwise_kind<erf>
ins(%A : memref<7x14x21xf32>) outs(%Out : memref<7x14x21xf32>)
+ linalg.elementwise kind=#linalg.elementwise_kind<sin>
+ ins(%A : memref<7x14x21xf32>) outs(%Out : memref<7x14x21xf32>)
+ linalg.elementwise kind=#linalg.elementwise_kind<cos>
+ ins(%A : memref<7x14x21xf32>) outs(%Out : memref<7x14x21xf32>)
+ linalg.elementwise kind=#linalg.elementwise_kind<tan>
+ ins(%A : memref<7x14x21xf32>) outs(%Out : memref<7x14x21xf32>)
return
}
@@ -77,6 +83,15 @@ func.func @unary_ops(%A: memref<7x14x21xf32>, %Out: memref<7x14x21xf32>) {
// CHECK: linalg.elementwise kind=#linalg.elementwise_kind<erf>
// CHECK-SAME: ins(%[[A]] : memref<7x14x21xf32>)
// CHECK-SAME: outs(%[[OUT]] : memref<7x14x21xf32>)
+// CHECK: linalg.elementwise kind=#linalg.elementwise_kind<sin>
+// CHECK-SAME: ins(%[[A]] : memref<7x14x21xf32>)
+// CHECK-SAME: outs(%[[OUT]] : memref<7x14x21xf32>)
+// CHECK: linalg.elementwise kind=#linalg.elementwise_kind<cos>
+// CHECK-SAME: ins(%[[A]] : memref<7x14x21xf32>)
+// CHECK-SAME: outs(%[[OUT]] : memref<7x14x21xf32>)
+// CHECK: linalg.elementwise kind=#linalg.elementwise_kind<tan>
+// CHECK-SAME: ins(%[[A]] : memref<7x14x21xf32>)
+// CHECK-SAME: outs(%[[OUT]] : memref<7x14x21xf32>)
// -----
diff --git a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
index 3d6c2962731c9..7b93c38259e20 100644
--- a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
+++ b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
@@ -100,7 +100,28 @@ func.func @unary_ops(%A: tensor<?x?x?xf32>, %Out: tensor<?x?x?xf32>) -> tensor<?
%v = math.erf %in : f32
linalg.yield %v : f32
} -> tensor<?x?x?xf32>
- return %12 : tensor<?x?x?xf32>
+ %13 = linalg.generic
+ {indexing_maps = [#umap, #umap], iterator_types = ["parallel", "parallel","parallel"]}
+ ins(%12 : tensor<?x?x?xf32>) outs(%Out : tensor<?x?x?xf32>) {
+ ^bb0(%in: f32, %out: f32):
+ %v = math.sin %in : f32
+ linalg.yield %v : f32
+ } -> tensor<?x?x?xf32>
+ %14 = linalg.generic
+ {indexing_maps = [#umap, #umap], iterator_types = ["parallel", "parallel","parallel"]}
+ ins(%13 : tensor<?x?x?xf32>) outs(%Out : tensor<?x?x?xf32>) {
+ ^bb0(%in: f32, %out: f32):
+ %v = math.cos %in : f32
+ linalg.yield %v : f32
+ } -> tensor<?x?x?xf32>
+ %15 = linalg.generic
+ {indexing_maps = [#umap, #umap], iterator_types = ["parallel", "parallel","parallel"]}
+ ins(%14 : tensor<?x?x?xf32>) outs(%Out : tensor<?x?x?xf32>) {
+ ^bb0(%in: f32, %out: f32):
+ %v = math.tan %in : f32
+ linalg.yield %v : f32
+ } -> tensor<?x?x?xf32>
+ return %15 : tensor<?x?x?xf32>
}
// ALL-LABEL: unary_ops
@@ -146,6 +167,15 @@ func.func @unary_ops(%A: tensor<?x?x?xf32>, %Out: tensor<?x?x?xf32>) -> tensor<?
// NAMED: %[[RES12:.+]] = linalg.erf
// NAMED-SAME: ins(%[[RES11]] : tensor<?x?x?xf32>)
// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// NAMED: %[[RES13:.+]] = linalg.sin
+// NAMED-SAME: ins(%[[RES12]] : tensor<?x?x?xf32>)
+// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// NAMED: %[[RES14:.+]] = linalg.cos
+// NAMED-SAME: ins(%[[RES13]] : tensor<?x?x?xf32>)
+// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// NAMED: %[[RES15:.+]] = linalg.tan
+// NAMED-SAME: ins(%[[RES14]] : tensor<?x?x?xf32>)
+// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
// CATEGORY: %[[RES0:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<exp>
// CATEGORY-SAME: ins(%[[A]] : tensor<?x?x?xf32>)
@@ -186,6 +216,15 @@ func.func @unary_ops(%A: tensor<?x?x?xf32>, %Out: tensor<?x?x?xf32>) -> tensor<?
// CATEGORY: %[[RES12:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<erf>
// CATEGORY-SAME: ins(%[[RES11]] : tensor<?x?x?xf32>)
// CATEGORY-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// CATEGORY: %[[RES13:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<sin>
+// CATEGORY-SAME: ins(%[[RES12]] : tensor<?x?x?xf32>)
+// CATEGORY-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// CATEGORY: %[[RES14:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<cos>
+// CATEGORY-SAME: ins(%[[RES13]] : tensor<?x?x?xf32>)
+// CATEGORY-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
+// CATEGORY: %[[RES15:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<tan>
+// CATEGORY-SAME: ins(%[[RES14]] : tensor<?x?x?xf32>)
+// CATEGORY-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
// -----
>From 5ffde541450c29122aa9a050c9ac91e927beecb6 Mon Sep 17 00:00:00 2001
From: Vinit Deodhar <vdeodhar at ah-vdeodhar-l.dhcp.mathworks.com>
Date: Wed, 3 Jun 2026 18:54:03 -0400
Subject: [PATCH 2/3] Remove sin, cos, tan named ops from YAML; use elementwise
only
---
.../Linalg/IR/LinalgNamedStructuredOps.yaml | 105 ------------------
.../Linalg/Transforms/NamedToElementwise.cpp | 6 -
.../Dialect/Linalg/Transforms/Specialize.cpp | 16 ++-
.../Dialect/Linalg/generalize-named-ops.mlir | 63 -----------
.../Linalg/specialize-generic-ops.mlir | 9 --
5 files changed, 10 insertions(+), 189 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
index f62ffe113d123..521afc991063f 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yaml
@@ -499,111 +499,6 @@ structured_op: !LinalgStructuredOpConfig
- !ScalarExpression
scalar_arg: I
--- !LinalgOpConfig
-metadata: !LinalgOpMetadata
- name: sin
- cpp_class_name: SinOp
- doc: |-
- Applies sin(x) elementwise.
-
- No numeric casting is performed on the input operand.
-structured_op: !LinalgStructuredOpConfig
- args:
- - !LinalgOperandDefConfig
- name: I
- 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<() -> ()>
- iterator_types: []
- assignments:
- - !ScalarAssign
- arg: O
- value: !ScalarExpression
- scalar_fn:
- kind: unary
- fn_name: sin
- operands:
- - !ScalarExpression
- scalar_arg: I
---- !LinalgOpConfig
-metadata: !LinalgOpMetadata
- name: cos
- cpp_class_name: CosOp
- doc: |-
- Applies cos(x) elementwise.
-
- No numeric casting is performed on the input operand.
-structured_op: !LinalgStructuredOpConfig
- args:
- - !LinalgOperandDefConfig
- name: I
- 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<() -> ()>
- iterator_types: []
- assignments:
- - !ScalarAssign
- arg: O
- value: !ScalarExpression
- scalar_fn:
- kind: unary
- fn_name: cos
- operands:
- - !ScalarExpression
- scalar_arg: I
---- !LinalgOpConfig
-metadata: !LinalgOpMetadata
- name: tan
- cpp_class_name: TanOp
- doc: |-
- Applies tan(x) elementwise.
-
- No numeric casting is performed on the input operand.
-structured_op: !LinalgStructuredOpConfig
- args:
- - !LinalgOperandDefConfig
- name: I
- 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<() -> ()>
- iterator_types: []
- assignments:
- - !ScalarAssign
- arg: O
- value: !ScalarExpression
- scalar_fn:
- kind: unary
- fn_name: tan
- operands:
- - !ScalarExpression
- scalar_arg: I
---- !LinalgOpConfig
metadata: !LinalgOpMetadata
name: add
cpp_class_name: AddOp
diff --git a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp b/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
index 02f6677cfbf7e..c9045566473cb 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/NamedToElementwise.cpp
@@ -48,9 +48,6 @@ ElementwiseKind getKind(Operation *op) {
.Case([](SquareOp) { return ElementwiseKind::square; })
.Case([](TanhOp) { return ElementwiseKind::tanh; })
.Case([](ErfOp) { return ElementwiseKind::erf; })
- .Case([](SinOp) { return ElementwiseKind::sin; })
- .Case([](CosOp) { return ElementwiseKind::cos; })
- .Case([](TanOp) { return ElementwiseKind::tan; })
.DefaultUnreachable("unhandled case in named to elementwise");
}
@@ -95,7 +92,4 @@ void mlir::linalg::populateLinalgNamedToElementwisePatterns(
patterns.add<NamedToElementwisePattern<SquareOp>>(patterns.getContext());
patterns.add<NamedToElementwisePattern<TanhOp>>(patterns.getContext());
patterns.add<NamedToElementwisePattern<ErfOp>>(patterns.getContext());
- patterns.add<NamedToElementwisePattern<SinOp>>(patterns.getContext());
- patterns.add<NamedToElementwisePattern<CosOp>>(patterns.getContext());
- patterns.add<NamedToElementwisePattern<TanOp>>(patterns.getContext());
}
diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 25f69974da29c..0ff488d24e4a7 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -223,12 +223,16 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
return replaceOp(TanhOp{}, ElementwiseKind::tanh);
if (isa<math::ErfOp>(op))
return replaceOp(ErfOp{}, ElementwiseKind::erf);
- if (isa<math::SinOp>(op))
- return replaceOp(SinOp{}, ElementwiseKind::sin);
- if (isa<math::CosOp>(op))
- return replaceOp(CosOp{}, ElementwiseKind::cos);
- if (isa<math::TanOp>(op))
- return replaceOp(TanOp{}, ElementwiseKind::tan);
+
+ // sin, cos, tan only have the category (elementwise) form, no named op.
+ if (emitCategoryOp) {
+ 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);
+ }
// At this point, we exhaustively checked the available unary named ops. The
// 1-input generic op might be representable as a `linalg.elementwise` that
diff --git a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
index 21d0d103774a6..e346bee901f1d 100644
--- a/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
+++ b/mlir/test/Dialect/Linalg/generalize-named-ops.mlir
@@ -806,69 +806,6 @@ func.func @generalize_erf(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>)
// -----
-func.func @generalize_sin(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
- linalg.sin ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
- return
-}
-
-// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-// CHECK: func @generalize_sin
-// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
-
-// CHECK: linalg.generic
-// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
-// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
-// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
-
-// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
-// CHECK-NEXT: %[[sin:.+]] = math.sin %[[BBARG0]] : f32
-// CHECK-NEXT: linalg.yield %[[sin]] : f32
-
-// -----
-
-func.func @generalize_cos(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
- linalg.cos ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
- return
-}
-
-// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-// CHECK: func @generalize_cos
-// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
-
-// CHECK: linalg.generic
-// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
-// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
-// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
-
-// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
-// CHECK-NEXT: %[[cos:.+]] = math.cos %[[BBARG0]] : f32
-// CHECK-NEXT: linalg.yield %[[cos]] : f32
-
-// -----
-
-func.func @generalize_tan(%arg: memref<7x14x21xf32>, %out: memref<7x14x21xf32>) {
- linalg.tan ins(%arg : memref<7x14x21xf32>) outs(%out : memref<7x14x21xf32>)
- return
-}
-
-// CHECK: #[[MAP:.+]] = affine_map<(d0, d1, d2) -> (d0, d1, d2)>
-
-// CHECK: func @generalize_tan
-// CHECK-SAME: (%[[ARG:.+]]: memref<7x14x21xf32>, %[[OUT:.+]]: memref<7x14x21xf32>)
-
-// CHECK: linalg.generic
-// CHECK-SAME: indexing_maps = [#[[MAP]], #[[MAP]]]
-// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel"]}
-// CHECK-SAME: ins(%[[LHS]] : memref<7x14x21xf32>) outs(%[[OUT]] : memref<7x14x21xf32>)
-
-// CHECK: ^{{.+}}(%[[BBARG0:.+]]: f32, %[[BBARG1:.+]]: f32)
-// CHECK-NEXT: %[[tan:.+]] = math.tan %[[BBARG0]] : f32
-// CHECK-NEXT: linalg.yield %[[tan]] : f32
-
-// -----
-
func.func @generalize_max(%lhs: memref<7x14x21xf32>, %rhs: memref<7x14x21xf32>,
%out: memref<7x14x21xf32>) {
linalg.max ins(%lhs, %rhs : memref<7x14x21xf32>, memref<7x14x21xf32>)
diff --git a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
index 7b93c38259e20..46193b4d7e992 100644
--- a/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
+++ b/mlir/test/Dialect/Linalg/specialize-generic-ops.mlir
@@ -167,15 +167,6 @@ func.func @unary_ops(%A: tensor<?x?x?xf32>, %Out: tensor<?x?x?xf32>) -> tensor<?
// NAMED: %[[RES12:.+]] = linalg.erf
// NAMED-SAME: ins(%[[RES11]] : tensor<?x?x?xf32>)
// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
-// NAMED: %[[RES13:.+]] = linalg.sin
-// NAMED-SAME: ins(%[[RES12]] : tensor<?x?x?xf32>)
-// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
-// NAMED: %[[RES14:.+]] = linalg.cos
-// NAMED-SAME: ins(%[[RES13]] : tensor<?x?x?xf32>)
-// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
-// NAMED: %[[RES15:.+]] = linalg.tan
-// NAMED-SAME: ins(%[[RES14]] : tensor<?x?x?xf32>)
-// NAMED-SAME: outs(%[[OUT]] : tensor<?x?x?xf32>) -> tensor<?x?x?xf32>
// CATEGORY: %[[RES0:.+]] = linalg.elementwise kind=#linalg.elementwise_kind<exp>
// CATEGORY-SAME: ins(%[[A]] : tensor<?x?x?xf32>)
>From 8344eda27b442b7590288d9cab1db7128741697b Mon Sep 17 00:00:00 2001
From: Vinit Deodhar <vdeodhar at ah-vdeodhar-l.dhcp.mathworks.com>
Date: Sat, 6 Jun 2026 03:12:41 -0400
Subject: [PATCH 3/3] Update comment explaining the specialization
---
mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp | 3 ++-
1 file changed, 2 insertions(+), 1 deletion(-)
diff --git a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
index 0ff488d24e4a7..fd95cfd5d3403 100644
--- a/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
+++ b/mlir/lib/Dialect/Linalg/Transforms/Specialize.cpp
@@ -224,7 +224,8 @@ static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
if (isa<math::ErfOp>(op))
return replaceOp(ErfOp{}, ElementwiseKind::erf);
- // sin, cos, tan only have the category (elementwise) form, no named op.
+ // sin, cos, tan only have the category (elementwise) form, but no
+ // linalg.* named op equivalent.
if (emitCategoryOp) {
if (isa<math::SinOp>(op))
return replaceOp(nullptr, ElementwiseKind::sin);
More information about the Mlir-commits
mailing list