[Mlir-commits] [mlir] 67d211a - [mlir][SPIR-V] Convert complex.neg and complex.conj in ComplexToSPIRV (#202898)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu Jun 11 01:48:07 PDT 2026
Author: Arseniy Obolenskiy
Date: 2026-06-11T10:48:02+02:00
New Revision: 67d211a220e79636cdef7667b1c429cb4fbd7660
URL: https://github.com/llvm/llvm-project/commit/67d211a220e79636cdef7667b1c429cb4fbd7660
DIFF: https://github.com/llvm/llvm-project/commit/67d211a220e79636cdef7667b1c429cb4fbd7660.diff
LOG: [mlir][SPIR-V] Convert complex.neg and complex.conj in ComplexToSPIRV (#202898)
Added:
Modified:
mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
Removed:
################################################################################
diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index b0069986b0b87..ba5fd06966153 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -186,6 +186,37 @@ struct AbsOpPattern final : OpConversionPattern<complex::AbsOp> {
}
};
+template <typename ComplexOp, bool NegateReal>
+struct NegationOpPattern 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 complexVal = adaptor.getComplex();
+
+ Value re =
+ spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
+ Value im =
+ spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
+
+ Value resultRe =
+ NegateReal ? spirv::FNegateOp::create(rewriter, loc, re) : re;
+ Value resultIm = spirv::FNegateOp::create(rewriter, loc, im);
+
+ rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+ op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+ return success();
+ }
+};
+
struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
using Base::Base;
@@ -236,6 +267,9 @@ 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,
+ NegationOpPattern<complex::NegOp, /*NegateReal=*/true>,
+ NegationOpPattern<complex::ConjOp, /*NegateReal=*/false>,
+ 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