[Mlir-commits] [mlir] [mlir][SPIR-V] Convert complex.neg and complex.conj in ComplexToSPIRV (PR #202898)

Arseniy Obolenskiy llvmlistbot at llvm.org
Wed Jun 10 02:14:33 PDT 2026


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

None

>From 2bdef4d1b3d9e968bd3f00d873d78d5d28b87bca Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 10 Jun 2026 11:13:11 +0200
Subject: [PATCH] [mlir][SPIR-V] Convert complex.neg and complex.conj in
 ComplexToSPIRV

---
 .../ComplexToSPIRV/ComplexToSPIRV.cpp         | 58 ++++++++++++++++++-
 .../ComplexToSPIRV/complex-to-spirv.mlir      | 31 ++++++++++
 2 files changed, 87 insertions(+), 2 deletions(-)

diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index b0069986b0b87..114552e40888e 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -186,6 +186,59 @@ struct AbsOpPattern final : OpConversionPattern<complex::AbsOp> {
   }
 };
 
+struct NegOpPattern final : OpConversionPattern<complex::NegOp> {
+  using Base::Base;
+
+  LogicalResult
+  matchAndRewrite(complex::NegOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
+    if (!spirvType)
+      return rewriter.notifyMatchFailure(op, "unable to convert result type");
+
+    Location loc = op.getLoc();
+    Value complexVal = adaptor.getComplex();
+
+    Value re =
+        spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
+    Value im =
+        spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
+
+    Value resultRe = spirv::FNegateOp::create(rewriter, loc, re);
+    Value resultIm = spirv::FNegateOp::create(rewriter, loc, im);
+
+    rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+        op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+    return success();
+  }
+};
+
+struct ConjOpPattern final : OpConversionPattern<complex::ConjOp> {
+  using Base::Base;
+
+  LogicalResult
+  matchAndRewrite(complex::ConjOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
+    if (!spirvType)
+      return rewriter.notifyMatchFailure(op, "unable to convert result type");
+
+    Location loc = op.getLoc();
+    Value complexVal = adaptor.getComplex();
+
+    Value re =
+        spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
+    Value im =
+        spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
+
+    Value resultIm = spirv::FNegateOp::create(rewriter, loc, im);
+
+    rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+        op, spirvType, llvm::ArrayRef<Value>{re, resultIm});
+    return success();
+  }
+};
+
 struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
   using Base::Base;
 
@@ -236,6 +289,7 @@ void mlir::populateComplexToSPIRVPatterns(
   patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern,
                ElementwiseBinaryOpPattern<complex::AddOp, spirv::FAddOp>,
                ElementwiseBinaryOpPattern<complex::SubOp, spirv::FSubOp>,
-               MulOpPattern, DivOpPattern, AbsOpPattern<spirv::GLSqrtOp>,
-               AbsOpPattern<spirv::CLSqrtOp>>(typeConverter, context);
+               MulOpPattern, DivOpPattern, NegOpPattern, ConjOpPattern,
+               AbsOpPattern<spirv::GLSqrtOp>, AbsOpPattern<spirv::CLSqrtOp>>(
+      typeConverter, context);
 }
diff --git a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
index bbf153feb1575..f818dea4ea375 100644
--- a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
+++ b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
@@ -132,6 +132,37 @@ func.func @complex_div(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
 
 // -----
 
+func.func @complex_neg(%arg: complex<f32>) -> complex<f32> {
+  %neg = complex.neg %arg : complex<f32>
+  return %neg : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_neg
+//  CHECK-SAME: %[[ARG:.+]]: complex<f32>
+//       CHECK:   %[[V:.+]] = builtin.unrealized_conversion_cast %[[ARG]] : complex<f32> to vector<2xf32>
+//       CHECK:   %[[RE:.+]] = spirv.CompositeExtract %[[V]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[IM:.+]] = spirv.CompositeExtract %[[V]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[NRE:.+]] = spirv.FNegate %[[RE]] : f32
+//       CHECK:   %[[NIM:.+]] = spirv.FNegate %[[IM]] : f32
+//       CHECK:   %[[CC:.+]] = spirv.CompositeConstruct %[[NRE]], %[[NIM]] : (f32, f32) -> vector<2xf32>
+
+// -----
+
+func.func @complex_conj(%arg: complex<f32>) -> complex<f32> {
+  %conj = complex.conj %arg : complex<f32>
+  return %conj : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_conj
+//  CHECK-SAME: %[[ARG:.+]]: complex<f32>
+//       CHECK:   %[[V:.+]] = builtin.unrealized_conversion_cast %[[ARG]] : complex<f32> to vector<2xf32>
+//       CHECK:   %[[RE:.+]] = spirv.CompositeExtract %[[V]][0 : i32] : vector<2xf32>
+//       CHECK:   %[[IM:.+]] = spirv.CompositeExtract %[[V]][1 : i32] : vector<2xf32>
+//       CHECK:   %[[NIM:.+]] = spirv.FNegate %[[IM]] : f32
+//       CHECK:   %[[CC:.+]] = spirv.CompositeConstruct %[[RE]], %[[NIM]] : (f32, f32) -> vector<2xf32>
+
+// -----
+
 func.func @complex_abs(%arg: complex<f32>) -> f32 {
   %abs = complex.abs %arg : complex<f32>
   return %abs : f32



More information about the Mlir-commits mailing list