[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