[Mlir-commits] [mlir] [mlir][complex] Add fold for 0+a -> a (PR #212199)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Tue Jul 28 08:23:43 PDT 2026


=?utf-8?q?Mattéo?= Rizza Murgier,=?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/4] [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/4] 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/4] [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();

>From 72f379bde64e426b8ce666472684589311f5864a 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 17:22:39 +0200
Subject: [PATCH 4/4] Revert bad fix

---
 mlir/lib/Dialect/Complex/IR/ComplexOps.cpp | 5 +++++
 1 file changed, 5 insertions(+)

diff --git a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
index 908837488eac1..b5323597b7ca4 100644
--- a/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
+++ b/mlir/lib/Dialect/Complex/IR/ComplexOps.cpp
@@ -256,6 +256,11 @@ 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