[Mlir-commits] [mlir] [mlir][SPIR-V] Add ComplexToSPIRV lowering for complex.sign (PR #216722)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Mon Aug 17 05:53:38 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
@llvm/pr-subscribers-mlir-spirv
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
Computes complex.sign as z / abs(z), using a component-wise spirv.Select to pick a zero result when z is zero without requiring SPIR-V 1.4 (a vector-result Select would need a scalar condition paired with a composite result, which is only legal from 1.4 on)
---
Full diff: https://github.com/llvm/llvm-project/pull/216722.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp (+57-1)
- (modified) mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir (+26)
``````````diff
diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index 884df06a4e705..e7bfb5548b24f 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -26,6 +26,13 @@ using namespace mlir;
namespace {
+/// Creates a scalar floating-point constant of the given value.
+Value createFPConstant(OpBuilder &builder, Location loc, Type type,
+ double value) {
+ return spirv::ConstantOp::create(
+ builder, loc, type, builder.getFloatAttr(cast<FloatType>(type), value));
+}
+
struct ConstantOpPattern final : OpConversionPattern<complex::ConstantOp> {
using Base::Base;
@@ -309,6 +316,54 @@ struct AngleOpPattern final : OpConversionPattern<complex::AngleOp> {
}
};
+/// Computes complex.sign as z / abs(z), selecting a zero result when z is
+/// zero to avoid a division by zero.
+template <typename SqrtOp>
+struct SignOpPattern final : OpConversionPattern<complex::SignOp> {
+ using OpConversionPattern<complex::SignOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(complex::SignOp 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 reSq = spirv::FMulOp::create(rewriter, loc, re, re);
+ Value imSq = spirv::FMulOp::create(rewriter, loc, im, im);
+ Value sum = spirv::FAddOp::create(rewriter, loc, reSq, imSq);
+ Value abs = SqrtOp::create(rewriter, loc, sum);
+
+ Value signRe = spirv::FDivOp::create(rewriter, loc, re, abs);
+ Value signIm = spirv::FDivOp::create(rewriter, loc, im, abs);
+
+ Value zero = createFPConstant(rewriter, loc, re.getType(), 0.0);
+ Value reIsZero = spirv::FOrdEqualOp::create(rewriter, loc, re, zero);
+ Value imIsZero = spirv::FOrdEqualOp::create(rewriter, loc, im, zero);
+ Value isZero =
+ spirv::LogicalAndOp::create(rewriter, loc, reIsZero, imIsZero);
+
+ // Select per component: spirv.Select with a scalar condition and a
+ // composite result requires SPIR-V 1.4, so operate on scalars here and
+ // assemble the result afterwards.
+ Value resultRe = spirv::SelectOp::create(rewriter, loc, isZero, re, signRe);
+ Value resultIm = spirv::SelectOp::create(rewriter, loc, isZero, im, signIm);
+
+ rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+ op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+ return success();
+ }
+};
+
} // namespace
//===----------------------------------------------------------------------===//
@@ -331,6 +386,7 @@ void mlir::populateComplexToSPIRVPatterns(
NegationOpPattern<complex::NegOp, /*NegateReal=*/true>,
NegationOpPattern<complex::ConjOp, /*NegateReal=*/false>,
AbsOpPattern<spirv::GLSqrtOp>, AbsOpPattern<spirv::CLSqrtOp>,
- AngleOpPattern<spirv::GLAtan2Op>, AngleOpPattern<spirv::CLAtan2Op>>(
+ AngleOpPattern<spirv::GLAtan2Op>, AngleOpPattern<spirv::CLAtan2Op>,
+ SignOpPattern<spirv::GLSqrtOp>, SignOpPattern<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 e060c1c0ad44a..10ee62d043e83 100644
--- a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
+++ b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
@@ -268,3 +268,29 @@ func.func @complex_angle_opencl(%arg: complex<f32>) -> f32 {
// CHECK: spirv.CL.atan2
}
+
+// -----
+
+func.func @complex_sign(%arg: complex<f32>) -> complex<f32> {
+ %sign = complex.sign %arg : complex<f32>
+ return %sign : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_sign
+// 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: %[[RESQ:.+]] = spirv.FMul %[[RE]], %[[RE]] : f32
+// CHECK: %[[IMSQ:.+]] = spirv.FMul %[[IM]], %[[IM]] : f32
+// CHECK: %[[SUM:.+]] = spirv.FAdd %[[RESQ]], %[[IMSQ]] : f32
+// CHECK: %[[ABS:.+]] = spirv.GL.Sqrt %[[SUM]] : f32
+// CHECK: %[[SIGNRE:.+]] = spirv.FDiv %[[RE]], %[[ABS]] : f32
+// CHECK: %[[SIGNIM:.+]] = spirv.FDiv %[[IM]], %[[ABS]] : f32
+// CHECK: %[[ZERO:.+]] = spirv.Constant 0.000000e+00 : f32
+// CHECK: %[[REZ:.+]] = spirv.FOrdEqual %[[RE]], %[[ZERO]] : f32
+// CHECK: %[[IMZ:.+]] = spirv.FOrdEqual %[[IM]], %[[ZERO]] : f32
+// CHECK: %[[ISZERO:.+]] = spirv.LogicalAnd %[[REZ]], %[[IMZ]] : i1
+// CHECK: %[[SELRE:.+]] = spirv.Select %[[ISZERO]], %[[RE]], %[[SIGNRE]] : i1, f32
+// CHECK: %[[SELIM:.+]] = spirv.Select %[[ISZERO]], %[[IM]], %[[SIGNIM]] : i1, f32
+// CHECK: %[[RESULT:.+]] = spirv.CompositeConstruct %[[SELRE]], %[[SELIM]] : (f32, f32) -> vector<2xf32>
``````````
</details>
https://github.com/llvm/llvm-project/pull/216722
More information about the Mlir-commits
mailing list