[Mlir-commits] [mlir] [mlir][complex] Add fold for 0+a -> a (PR #212199)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue Jul 28 05:29:15 PDT 2026
=?utf-8?q?Mattéo?= Rizza Murgier,=?utf-8?q?Mattéo?= Rizza Murgier
Message-ID:
In-Reply-To: <llvm.org/llvm/llvm-project/pull/212199 at github.com>
https://github.com/Brythzz updated https://github.com/llvm/llvm-project/pull/212199
>From fbabd064289db2e19b565420106daa1133e98a13 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?Matt=C3=A9o=20Rizza=20Murgier?=
<matteo.rizza-murgier at sipearl.com>
Date: Mon, 20 Jul 2026 12:49:36 +0200
Subject: [PATCH 1/3] [mlir][complex] Add fold for 0+a -> a
---
mlir/lib/Dialect/Complex/IR/ComplexOps.cpp | 9 +++++++++
mlir/test/Dialect/Complex/canonicalize.mlir | 14 ++++++++++++--
2 files changed, 21 insertions(+), 2 deletions(-)
diff --git a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
index b5323597b7ca4..d52f7fc316d09 100644
--- a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
+++ b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
@@ -270,6 +270,15 @@ OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
}
}
+ // complex.add(complex.constant<0.0, 0.0>, a) -> a
+ if (auto constantOp = getLhs().getDefiningOp<ConstantOp>()) {
+ auto arrayAttr = constantOp.getValue();
+ if (llvm::cast<FloatAttr>(arrayAttr[0]).getValue().isZero() &&
+ llvm::cast<FloatAttr>(arrayAttr[1]).getValue().isZero()) {
+ return getRhs();
+ }
+ }
+
return {};
}
diff --git a/mlir/test/Dialect/Complex/canonicalize.mlir b/mlir/test/Dialect/Complex/canonicalize.mlir
index 1c5216c82e5c3..8ee4c62b216a2 100644
--- a/mlir/test/Dialect/Complex/canonicalize.mlir
+++ b/mlir/test/Dialect/Complex/canonicalize.mlir
@@ -125,8 +125,8 @@ func.func @complex_conj_conj() -> complex<f32> {
return %conj2 : complex<f32>
}
-// CHECK-LABEL: func @complex_add_zero
-func.func @complex_add_zero() -> complex<f32> {
+// CHECK-LABEL: func @complex_add_zero_rhs
+func.func @complex_add_zero_rhs() -> complex<f32> {
%complex1 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
%complex2 = complex.constant [0.0 : f32, 0.0 : f32] : complex<f32>
// CHECK: %[[CPLX:.*]] = complex.constant [1.000000e+00 : f32, 0.000000e+00 : f32] : complex<f32>
@@ -135,6 +135,16 @@ func.func @complex_add_zero() -> complex<f32> {
return %add : complex<f32>
}
+// CHECK-LABEL: func @complex_add_zero_lhs
+func.func @complex_add_zero_lhs() -> complex<f32> {
+ %complex1 = complex.constant [0.0 : f32, 0.0 : f32] : complex<f32>
+ %complex2 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
+ // CHECK: %[[CPLX:.*]] = complex.constant [1.000000e+00 : f32, 0.000000e+00 : f32] : complex<f32>
+ // CHECK-NEXT: return %[[CPLX:.*]] : complex<f32>
+ %add = complex.add %complex1, %complex2 : complex<f32>
+ return %add : complex<f32>
+}
+
// CHECK-LABEL: func @complex_sub_add_lhs
func.func @complex_sub_add_lhs() -> complex<f32> {
%complex1 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
>From d9297b489592bc68082b38c65abc04c6cbda687a Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?Matt=C3=A9o=20Rizza=20Murgier?=
<matteo.rizza-murgier at sipearl.com>
Date: Tue, 28 Jul 2026 14:22:02 +0200
Subject: [PATCH 2/3] Revert "[mlir][complex] Add fold for 0+a -> a"
This reverts commit fbabd064289db2e19b565420106daa1133e98a13.
---
mlir/lib/Dialect/Complex/IR/ComplexOps.cpp | 9 ---------
mlir/test/Dialect/Complex/canonicalize.mlir | 14 ++------------
2 files changed, 2 insertions(+), 21 deletions(-)
diff --git a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
index d52f7fc316d09..b5323597b7ca4 100644
--- a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
+++ b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
@@ -270,15 +270,6 @@ OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
}
}
- // complex.add(complex.constant<0.0, 0.0>, a) -> a
- if (auto constantOp = getLhs().getDefiningOp<ConstantOp>()) {
- auto arrayAttr = constantOp.getValue();
- if (llvm::cast<FloatAttr>(arrayAttr[0]).getValue().isZero() &&
- llvm::cast<FloatAttr>(arrayAttr[1]).getValue().isZero()) {
- return getRhs();
- }
- }
-
return {};
}
diff --git a/mlir/test/Dialect/Complex/canonicalize.mlir b/mlir/test/Dialect/Complex/canonicalize.mlir
index 8ee4c62b216a2..1c5216c82e5c3 100644
--- a/mlir/test/Dialect/Complex/canonicalize.mlir
+++ b/mlir/test/Dialect/Complex/canonicalize.mlir
@@ -125,8 +125,8 @@ func.func @complex_conj_conj() -> complex<f32> {
return %conj2 : complex<f32>
}
-// CHECK-LABEL: func @complex_add_zero_rhs
-func.func @complex_add_zero_rhs() -> complex<f32> {
+// CHECK-LABEL: func @complex_add_zero
+func.func @complex_add_zero() -> complex<f32> {
%complex1 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
%complex2 = complex.constant [0.0 : f32, 0.0 : f32] : complex<f32>
// CHECK: %[[CPLX:.*]] = complex.constant [1.000000e+00 : f32, 0.000000e+00 : f32] : complex<f32>
@@ -135,16 +135,6 @@ func.func @complex_add_zero_rhs() -> complex<f32> {
return %add : complex<f32>
}
-// CHECK-LABEL: func @complex_add_zero_lhs
-func.func @complex_add_zero_lhs() -> complex<f32> {
- %complex1 = complex.constant [0.0 : f32, 0.0 : f32] : complex<f32>
- %complex2 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
- // CHECK: %[[CPLX:.*]] = complex.constant [1.000000e+00 : f32, 0.000000e+00 : f32] : complex<f32>
- // CHECK-NEXT: return %[[CPLX:.*]] : complex<f32>
- %add = complex.add %complex1, %complex2 : complex<f32>
- return %add : complex<f32>
-}
-
// CHECK-LABEL: func @complex_sub_add_lhs
func.func @complex_sub_add_lhs() -> complex<f32> {
%complex1 = complex.constant [1.0 : f32, 0.0 : f32] : complex<f32>
>From 88bd2cacebff3231fa0d55f0ccf2fbb7ab7b6b72 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?Matt=C3=A9o=20Rizza=20Murgier?=
<matteo.rizza-murgier at sipearl.com>
Date: Tue, 28 Jul 2026 14:26:25 +0200
Subject: [PATCH 3/3] [mlir][complex] Make complex AddOp commutative
Also remove the redundant folding pattern b+(a-b)->a
---
mlir/include/mlir/Dialect/Complex/IR/ComplexOps.td | 2 +-
mlir/lib/Dialect/Complex/IR/ComplexOps.cpp | 5 -----
2 files changed, 1 insertion(+), 6 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Complex/IR/ComplexOps.td b/mlir/include/mlir/Dialect/Complex/IR/ComplexOps.td
index 828379ded14b3..c6a0b8e9dfeea 100644
--- a/mlir/include/mlir/Dialect/Complex/IR/ComplexOps.td
+++ b/mlir/include/mlir/Dialect/Complex/IR/ComplexOps.td
@@ -66,7 +66,7 @@ def AbsOp : ComplexUnaryOp<"abs",
// AddOp
//===----------------------------------------------------------------------===//
-def AddOp : ComplexArithmeticOp<"add"> {
+def AddOp : ComplexArithmeticOp<"add", [Commutative]> {
let summary = "complex addition";
let description = [{
The `add` operation takes two complex numbers and returns their sum.
diff --git a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
index b5323597b7ca4..908837488eac1 100644
--- a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
+++ b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
@@ -256,11 +256,6 @@ OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
if (getRhs() == sub.getRhs())
return sub.getLhs();
- // complex.add(b, complex.sub(a, b)) -> a
- if (auto sub = getRhs().getDefiningOp<SubOp>())
- if (getLhs() == sub.getRhs())
- return sub.getLhs();
-
// complex.add(a, complex.constant<0.0, 0.0>) -> a
if (auto constantOp = getRhs().getDefiningOp<ConstantOp>()) {
auto arrayAttr = constantOp.getValue();
More information about the Mlir-commits
mailing list