[Mlir-commits] [mlir] [mlir][SPIR-V] Convert complex.abs (PR #202026)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Sat Jun 6 03:20:10 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
---
Full diff: https://github.com/llvm/llvm-project/pull/202026.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp (+31-1)
- (modified) mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir (+37)
``````````diff
diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index 2e8eb88f2bfc1..b0069986b0b87 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -157,6 +157,35 @@ struct MulOpPattern final : OpConversionPattern<complex::MulOp> {
}
};
+template <typename SqrtOp>
+struct AbsOpPattern final : OpConversionPattern<complex::AbsOp> {
+ using OpConversionPattern<complex::AbsOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(complex::AbsOp 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);
+
+ rewriter.replaceOpWithNewOp<SqrtOp>(op, sum);
+ return success();
+ }
+};
+
struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
using Base::Base;
@@ -207,5 +236,6 @@ void mlir::populateComplexToSPIRVPatterns(
patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern,
ElementwiseBinaryOpPattern<complex::AddOp, spirv::FAddOp>,
ElementwiseBinaryOpPattern<complex::SubOp, spirv::FSubOp>,
- MulOpPattern, DivOpPattern>(typeConverter, context);
+ MulOpPattern, DivOpPattern, 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 69aa9765bb6f2..bbf153feb1575 100644
--- a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
+++ b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
@@ -129,3 +129,40 @@ func.func @complex_div(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
// CHECK: %[[RE:.+]] = spirv.FDiv %[[NRE]], %[[DENOM]] : f32
// CHECK: %[[IM:.+]] = spirv.FDiv %[[NIM]], %[[DENOM]] : f32
// CHECK: %[[CC:.+]] = spirv.CompositeConstruct %[[RE]], %[[IM]] : (f32, f32) -> vector<2xf32>
+
+// -----
+
+func.func @complex_abs(%arg: complex<f32>) -> f32 {
+ %abs = complex.abs %arg : complex<f32>
+ return %abs : f32
+}
+
+// CHECK-LABEL: func.func @complex_abs
+// 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: return %[[ABS]] : f32
+
+// -----
+
+module attributes {
+ spirv.target_env = #spirv.target_env<#spirv.vce<v1.0, [Kernel], []>, #spirv.resource_limits<>>
+} {
+
+func.func @complex_abs_opencl(%arg: complex<f32>) -> f32 {
+ %abs = complex.abs %arg : complex<f32>
+ return %abs : f32
+}
+
+// CHECK-LABEL: func.func @complex_abs_opencl
+// CHECK: spirv.FMul
+// CHECK: spirv.FMul
+// CHECK: %[[SUM:.+]] = spirv.FAdd
+// CHECK: spirv.CL.sqrt %[[SUM]] : f32
+
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/202026
More information about the Mlir-commits
mailing list