[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