[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