[Mlir-commits] [mlir] [mlir][SPIR-V][complex] Convert complex.add/sub/mul/div to SPIR-V ops (PR #200123)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu May 28 00:34:11 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/200123.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp (+104-2)
- (modified) mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir (+82)
``````````diff
diff --git a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
index a2a2faf8ea570..2e8eb88f2bfc1 100644
--- a/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
+++ b/mlir/lib/Conversion/ComplexToSPIRV/ComplexToSPIRV.cpp
@@ -94,6 +94,106 @@ struct ImOpPattern final : OpConversionPattern<complex::ImOp> {
}
};
+template <typename ComplexOp, typename SPIRVOp>
+struct ElementwiseBinaryOpPattern 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 lhs = adaptor.getLhs();
+ Value rhs = adaptor.getRhs();
+
+ Value lhsRe = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
+ Value lhsIm = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
+ Value rhsRe = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
+ Value rhsIm = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
+
+ Value resultRe = SPIRVOp::create(rewriter, loc, lhsRe, rhsRe);
+ Value resultIm = SPIRVOp::create(rewriter, loc, lhsIm, rhsIm);
+
+ rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+ op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+ return success();
+ }
+};
+
+struct MulOpPattern final : OpConversionPattern<complex::MulOp> {
+ using Base::Base;
+
+ LogicalResult
+ matchAndRewrite(complex::MulOp 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 lhs = adaptor.getLhs();
+ Value rhs = adaptor.getRhs();
+
+ Value a = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
+ Value b = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
+ Value c = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
+ Value d = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
+
+ Value ac = spirv::FMulOp::create(rewriter, loc, a, c);
+ Value bd = spirv::FMulOp::create(rewriter, loc, b, d);
+ Value ad = spirv::FMulOp::create(rewriter, loc, a, d);
+ Value bc = spirv::FMulOp::create(rewriter, loc, b, c);
+ Value resultRe = spirv::FSubOp::create(rewriter, loc, ac, bd);
+ Value resultIm = spirv::FAddOp::create(rewriter, loc, ad, bc);
+
+ rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+ op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+ return success();
+ }
+};
+
+struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
+ using Base::Base;
+
+ LogicalResult
+ matchAndRewrite(complex::DivOp 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 lhs = adaptor.getLhs();
+ Value rhs = adaptor.getRhs();
+
+ Value a = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
+ Value b = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
+ Value c = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
+ Value d = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
+
+ Value ac = spirv::FMulOp::create(rewriter, loc, a, c);
+ Value bd = spirv::FMulOp::create(rewriter, loc, b, d);
+ Value bc = spirv::FMulOp::create(rewriter, loc, b, c);
+ Value ad = spirv::FMulOp::create(rewriter, loc, a, d);
+ Value cc = spirv::FMulOp::create(rewriter, loc, c, c);
+ Value dd = spirv::FMulOp::create(rewriter, loc, d, d);
+ Value denom = spirv::FAddOp::create(rewriter, loc, cc, dd);
+ Value numRe = spirv::FAddOp::create(rewriter, loc, ac, bd);
+ Value numIm = spirv::FSubOp::create(rewriter, loc, bc, ad);
+ Value resultRe = spirv::FDivOp::create(rewriter, loc, numRe, denom);
+ Value resultIm = spirv::FDivOp::create(rewriter, loc, numIm, denom);
+
+ rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
+ op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
+ return success();
+ }
+};
+
} // namespace
//===----------------------------------------------------------------------===//
@@ -104,6 +204,8 @@ void mlir::populateComplexToSPIRVPatterns(
const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns) {
MLIRContext *context = patterns.getContext();
- patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern>(
- typeConverter, context);
+ patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern,
+ ElementwiseBinaryOpPattern<complex::AddOp, spirv::FAddOp>,
+ ElementwiseBinaryOpPattern<complex::SubOp, spirv::FSubOp>,
+ MulOpPattern, DivOpPattern>(typeConverter, context);
}
diff --git a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
index 45f38d435c50b..69aa9765bb6f2 100644
--- a/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
+++ b/mlir/test/Conversion/ComplexToSPIRV/complex-to-spirv.mlir
@@ -47,3 +47,85 @@ func.func @complex_const() -> complex<f32> {
// CHECK-LABEL: func.func @complex_const()
// CHECK: spirv.Constant dense<[0x7FC00000, 0.000000e+00]> : vector<2xf32>
+
+// -----
+
+func.func @complex_add(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
+ %0 = complex.add %lhs, %rhs : complex<f32>
+ return %0 : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_add
+// CHECK-SAME: (%[[LHS:.+]]: complex<f32>, %[[RHS:.+]]: complex<f32>)
+// CHECK-DAG: %[[LV:.+]] = builtin.unrealized_conversion_cast %[[LHS]] : complex<f32> to vector<2xf32>
+// CHECK-DAG: %[[RV:.+]] = builtin.unrealized_conversion_cast %[[RHS]] : complex<f32> to vector<2xf32>
+// CHECK: %[[LRE:.+]] = spirv.CompositeExtract %[[LV]][0 : i32] : vector<2xf32>
+// CHECK: %[[LIM:.+]] = spirv.CompositeExtract %[[LV]][1 : i32] : vector<2xf32>
+// CHECK: %[[RRE:.+]] = spirv.CompositeExtract %[[RV]][0 : i32] : vector<2xf32>
+// CHECK: %[[RIM:.+]] = spirv.CompositeExtract %[[RV]][1 : i32] : vector<2xf32>
+// CHECK: %[[RE:.+]] = spirv.FAdd %[[LRE]], %[[RRE]] : f32
+// CHECK: %[[IM:.+]] = spirv.FAdd %[[LIM]], %[[RIM]] : f32
+// CHECK: %[[CC:.+]] = spirv.CompositeConstruct %[[RE]], %[[IM]] : (f32, f32) -> vector<2xf32>
+
+// -----
+
+func.func @complex_sub(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
+ %0 = complex.sub %lhs, %rhs : complex<f32>
+ return %0 : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_sub
+// CHECK: spirv.FSub
+// CHECK: spirv.FSub
+// CHECK: spirv.CompositeConstruct
+
+// -----
+
+func.func @complex_mul(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
+ %0 = complex.mul %lhs, %rhs : complex<f32>
+ return %0 : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_mul
+// CHECK-SAME: (%[[LHS:.+]]: complex<f32>, %[[RHS:.+]]: complex<f32>)
+// CHECK-DAG: %[[LV:.+]] = builtin.unrealized_conversion_cast %[[LHS]] : complex<f32> to vector<2xf32>
+// CHECK-DAG: %[[RV:.+]] = builtin.unrealized_conversion_cast %[[RHS]] : complex<f32> to vector<2xf32>
+// CHECK: %[[A:.+]] = spirv.CompositeExtract %[[LV]][0 : i32] : vector<2xf32>
+// CHECK: %[[B:.+]] = spirv.CompositeExtract %[[LV]][1 : i32] : vector<2xf32>
+// CHECK: %[[C:.+]] = spirv.CompositeExtract %[[RV]][0 : i32] : vector<2xf32>
+// CHECK: %[[D:.+]] = spirv.CompositeExtract %[[RV]][1 : i32] : vector<2xf32>
+// CHECK: %[[AC:.+]] = spirv.FMul %[[A]], %[[C]] : f32
+// CHECK: %[[BD:.+]] = spirv.FMul %[[B]], %[[D]] : f32
+// CHECK: %[[AD:.+]] = spirv.FMul %[[A]], %[[D]] : f32
+// CHECK: %[[BC:.+]] = spirv.FMul %[[B]], %[[C]] : f32
+// CHECK: %[[RE:.+]] = spirv.FSub %[[AC]], %[[BD]] : f32
+// CHECK: %[[IM:.+]] = spirv.FAdd %[[AD]], %[[BC]] : f32
+// CHECK: %[[CC:.+]] = spirv.CompositeConstruct %[[RE]], %[[IM]] : (f32, f32) -> vector<2xf32>
+
+// -----
+
+func.func @complex_div(%lhs: complex<f32>, %rhs: complex<f32>) -> complex<f32> {
+ %0 = complex.div %lhs, %rhs : complex<f32>
+ return %0 : complex<f32>
+}
+
+// CHECK-LABEL: func.func @complex_div
+// CHECK-SAME: (%[[LHS:.+]]: complex<f32>, %[[RHS:.+]]: complex<f32>)
+// CHECK-DAG: %[[LV:.+]] = builtin.unrealized_conversion_cast %[[LHS]] : complex<f32> to vector<2xf32>
+// CHECK-DAG: %[[RV:.+]] = builtin.unrealized_conversion_cast %[[RHS]] : complex<f32> to vector<2xf32>
+// CHECK: %[[A:.+]] = spirv.CompositeExtract %[[LV]][0 : i32] : vector<2xf32>
+// CHECK: %[[B:.+]] = spirv.CompositeExtract %[[LV]][1 : i32] : vector<2xf32>
+// CHECK: %[[C:.+]] = spirv.CompositeExtract %[[RV]][0 : i32] : vector<2xf32>
+// CHECK: %[[D:.+]] = spirv.CompositeExtract %[[RV]][1 : i32] : vector<2xf32>
+// CHECK: %[[AC:.+]] = spirv.FMul %[[A]], %[[C]] : f32
+// CHECK: %[[BD:.+]] = spirv.FMul %[[B]], %[[D]] : f32
+// CHECK: %[[BC:.+]] = spirv.FMul %[[B]], %[[C]] : f32
+// CHECK: %[[AD:.+]] = spirv.FMul %[[A]], %[[D]] : f32
+// CHECK: %[[CC2:.+]] = spirv.FMul %[[C]], %[[C]] : f32
+// CHECK: %[[DD:.+]] = spirv.FMul %[[D]], %[[D]] : f32
+// CHECK: %[[DENOM:.+]] = spirv.FAdd %[[CC2]], %[[DD]] : f32
+// CHECK: %[[NRE:.+]] = spirv.FAdd %[[AC]], %[[BD]] : f32
+// CHECK: %[[NIM:.+]] = spirv.FSub %[[BC]], %[[AD]] : 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>
``````````
</details>
https://github.com/llvm/llvm-project/pull/200123
More information about the Mlir-commits
mailing list