[Mlir-commits] [mlir] [mlir][ComplexToSPIRV] Add lowering for complex.eq and complex.neq (PR #206279)

Arseniy Obolenskiy llvmlistbot at llvm.org
Sat Jun 27 12:23:04 PDT 2026


https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/206279

None

>From bad28b6126511d13d96538295cc492f7e590a36c Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Sat, 27 Jun 2026 21:19:04 +0200
Subject: [PATCH] [mlir][ComplexToSPIRV] Add lowering for complex.eq and
 complex.neq

---
 .../ComplexToSPIRV/ComplexToSPIRV.cpp         | 34 ++++++++++++++++
 .../ComplexToSPIRV/complex-to-spirv.mlir      | 40 +++++++++++++++++++
 2 files changed, 74 insertions(+)

diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index ba5fd06966153..eb0f81dfecf30 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -125,6 +125,36 @@ struct ElementwiseBinaryOpPattern final : OpConversionPattern<ComplexOp> {
   }
 };
 
+template <typename ComplexOp, typename SPIRVCompareOp, typename SPIRVCombinerOp>
+struct ComparisonOpPattern final : OpConversionPattern<ComplexOp> {
+  using OpConversionPattern<ComplexOp>::OpConversionPattern;
+  using OpAdaptor = typename ComplexOp::Adaptor;
+
+  LogicalResult
+  matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type spirvType =
+        this->getTypeConverter()->convertType(op.getResult().getType());
+    if (!spirvType)
+      return rewriter.notifyMatchFailure(op, "unable to convert result type");
+
+    Location loc = op.getLoc();
+    Value lhs = adaptor.getLhs();
+    Value rhs = adaptor.getRhs();
+
+    Value lhsRe = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
+    Value lhsIm = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
+    Value rhsRe = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
+    Value rhsIm = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
+
+    Value cmpRe = SPIRVCompareOp::create(rewriter, loc, lhsRe, rhsRe);
+    Value cmpIm = SPIRVCompareOp::create(rewriter, loc, lhsIm, rhsIm);
+
+    rewriter.replaceOpWithNewOp<SPIRVCombinerOp>(op, spirvType, cmpRe, cmpIm);
+    return success();
+  }
+};
+
 struct MulOpPattern final : OpConversionPattern<complex::MulOp> {
   using Base::Base;
 
@@ -267,6 +297,10 @@ void mlir::populateComplexToSPIRVPatterns(
   patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern,
                ElementwiseBinaryOpPattern<complex::AddOp, spirv::FAddOp>,
                ElementwiseBinaryOpPattern<complex::SubOp, spirv::FSubOp>,
+               ComparisonOpPattern<complex::EqualOp, spirv::FOrdEqualOp,
+                                   spirv::LogicalAndOp>,
+               ComparisonOpPattern<complex::NotEqualOp, spirv::FUnordNotEqualOp,
+                                   spirv::LogicalOrOp>,
                MulOpPattern, DivOpPattern,
                NegationOpPattern<complex::NegOp, /*NegateReal=*/true>,
                NegationOpPattern<complex::ConjOp, /*NegateReal=*/false>,
diff --git a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
index f818dea4ea375..deb4eb6d9d08c 100644
--- a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
+++ b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
@@ -132,6 +132,46 @@ func.func @complex_div(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
 
 // -----
 
+func.func @complex_eq(%lhs: complex<f32>, %rhs: complex<f32>) -> i1 {
+  %0 = complex.eq %lhs, %rhs : complex<f32>
+  return %0 : i1
+}
+
+// CHECK-LABEL: func.func @complex_eq
+//  CHECK-SAME: (%[[LHS:.+]]: complex<f32>, %[[RHS:.+]]: complex<f32>)
+//   CHECK-DAG:   %[[LV:.+]] = builtin.unrealized_conversion_cast %[[LHS]] : complex<f32> to vector<2xf32>
+//   CHECK-DAG:   %[[RV:.+]] = builtin.unrealized_conversion_cast %[[RHS]] : complex<f32> to vector<2xf32>
+//       CHECK:   %[[LRE:.+]] = spirv.CompositeExtract %[[LV]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[LIM:.+]] = spirv.CompositeExtract %[[LV]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[RRE:.+]] = spirv.CompositeExtract %[[RV]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[RIM:.+]] = spirv.CompositeExtract %[[RV]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[REQ:.+]] = spirv.FOrdEqual %[[LRE]], %[[RRE]] : f32
+//       CHECK:   %[[IMEQ:.+]] = spirv.FOrdEqual %[[LIM]], %[[RIM]] : f32
+//       CHECK:   %[[EQ:.+]] = spirv.LogicalAnd %[[REQ]], %[[IMEQ]] : i1
+//       CHECK:   return %[[EQ]] : i1
+
+// -----
+
+func.func @complex_neq(%lhs: complex<f32>, %rhs: complex<f32>) -> i1 {
+  %0 = complex.neq %lhs, %rhs : complex<f32>
+  return %0 : i1
+}
+
+// CHECK-LABEL: func.func @complex_neq
+//  CHECK-SAME: (%[[LHS:.+]]: complex<f32>, %[[RHS:.+]]: complex<f32>)
+//   CHECK-DAG:   %[[LV:.+]] = builtin.unrealized_conversion_cast %[[LHS]] : complex<f32> to vector<2xf32>
+//   CHECK-DAG:   %[[RV:.+]] = builtin.unrealized_conversion_cast %[[RHS]] : complex<f32> to vector<2xf32>
+//       CHECK:   %[[LRE:.+]] = spirv.CompositeExtract %[[LV]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[LIM:.+]] = spirv.CompositeExtract %[[LV]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[RRE:.+]] = spirv.CompositeExtract %[[RV]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[RIM:.+]] = spirv.CompositeExtract %[[RV]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[RNE:.+]] = spirv.FUnordNotEqual %[[LRE]], %[[RRE]] : f32
+//       CHECK:   %[[IMNE:.+]] = spirv.FUnordNotEqual %[[LIM]], %[[RIM]] : f32
+//       CHECK:   %[[NE:.+]] = spirv.LogicalOr %[[RNE]], %[[IMNE]] : i1
+//       CHECK:   return %[[NE]] : i1
+
+// -----
+
 func.func @complex_neg(%arg: complex<f32>) -> complex<f32> {
   %neg = complex.neg %arg : complex<f32>
   return %neg : complex<f32>



More information about the Mlir-commits mailing list