MLIR 24.0.0git
ComplexToSPIRV.cpp
Go to the documentation of this file.
1//===- ComplexToSPIRV.cpp - Complex to SPIR-V Patterns --------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements patterns to convert Complex dialect to SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
18
19#define DEBUG_TYPE "complex-to-spirv-pattern"
20
21using namespace mlir;
22
23//===----------------------------------------------------------------------===//
24// Operation conversion
25//===----------------------------------------------------------------------===//
26
27namespace {
28
29struct ConstantOpPattern final : OpConversionPattern<complex::ConstantOp> {
30 using Base::Base;
31
32 LogicalResult
33 matchAndRewrite(complex::ConstantOp constOp, OpAdaptor adaptor,
34 ConversionPatternRewriter &rewriter) const override {
35 auto spirvType =
36 getTypeConverter()->convertType<ShapedType>(constOp.getType());
37 if (!spirvType)
38 return rewriter.notifyMatchFailure(constOp,
39 "unable to convert result type");
40
41 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
42 constOp, spirvType,
43 DenseElementsAttr::get(spirvType, constOp.getValue().getValue()));
44 return success();
45 }
46};
47
48struct CreateOpPattern final : OpConversionPattern<complex::CreateOp> {
49 using Base::Base;
50
51 LogicalResult
52 matchAndRewrite(complex::CreateOp createOp, OpAdaptor adaptor,
53 ConversionPatternRewriter &rewriter) const override {
54 Type spirvType = getTypeConverter()->convertType(createOp.getType());
55 if (!spirvType)
56 return rewriter.notifyMatchFailure(createOp,
57 "unable to convert result type");
58
59 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
60 createOp, spirvType, adaptor.getOperands());
61 return success();
62 }
63};
64
65struct ReOpPattern final : OpConversionPattern<complex::ReOp> {
66 using Base::Base;
67
68 LogicalResult
69 matchAndRewrite(complex::ReOp reOp, OpAdaptor adaptor,
70 ConversionPatternRewriter &rewriter) const override {
71 Type spirvType = getTypeConverter()->convertType(reOp.getType());
72 if (!spirvType)
73 return rewriter.notifyMatchFailure(reOp, "unable to convert result type");
74
75 rewriter.replaceOpWithNewOp<spirv::CompositeExtractOp>(
76 reOp, adaptor.getComplex(), llvm::ArrayRef(0));
77 return success();
78 }
79};
80
81struct ImOpPattern final : OpConversionPattern<complex::ImOp> {
82 using Base::Base;
83
84 LogicalResult
85 matchAndRewrite(complex::ImOp imOp, OpAdaptor adaptor,
86 ConversionPatternRewriter &rewriter) const override {
87 Type spirvType = getTypeConverter()->convertType(imOp.getType());
88 if (!spirvType)
89 return rewriter.notifyMatchFailure(imOp, "unable to convert result type");
90
91 rewriter.replaceOpWithNewOp<spirv::CompositeExtractOp>(
92 imOp, adaptor.getComplex(), llvm::ArrayRef(1));
93 return success();
94 }
95};
96
97template <typename ComplexOp, typename SPIRVOp>
98struct ElementwiseBinaryOpPattern final : OpConversionPattern<ComplexOp> {
99 using OpConversionPattern<ComplexOp>::OpConversionPattern;
100 using OpAdaptor = typename ComplexOp::Adaptor;
101
102 LogicalResult
103 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
104 ConversionPatternRewriter &rewriter) const override {
105 Type spirvType =
106 this->getTypeConverter()->convertType(op.getResult().getType());
107 if (!spirvType)
108 return rewriter.notifyMatchFailure(op, "unable to convert result type");
109
110 Location loc = op.getLoc();
111 Value lhs = adaptor.getLhs();
112 Value rhs = adaptor.getRhs();
113
114 Value lhsRe = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
115 Value lhsIm = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
116 Value rhsRe = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
117 Value rhsIm = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
118
119 Value resultRe = SPIRVOp::create(rewriter, loc, lhsRe, rhsRe);
120 Value resultIm = SPIRVOp::create(rewriter, loc, lhsIm, rhsIm);
121
122 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
123 op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
124 return success();
125 }
126};
127
128template <typename ComplexOp, typename SPIRVCompareOp, typename SPIRVCombinerOp>
129struct ComparisonOpPattern final : OpConversionPattern<ComplexOp> {
130 using OpConversionPattern<ComplexOp>::OpConversionPattern;
131 using OpAdaptor = typename ComplexOp::Adaptor;
132
133 LogicalResult
134 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
135 ConversionPatternRewriter &rewriter) const override {
136 Type spirvType =
137 this->getTypeConverter()->convertType(op.getResult().getType());
138 if (!spirvType)
139 return rewriter.notifyMatchFailure(op, "unable to convert result type");
140
141 Location loc = op.getLoc();
142 Value lhs = adaptor.getLhs();
143 Value rhs = adaptor.getRhs();
144
145 Value lhsRe = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
146 Value lhsIm = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
147 Value rhsRe = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
148 Value rhsIm = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
149
150 Value cmpRe = SPIRVCompareOp::create(rewriter, loc, lhsRe, rhsRe);
151 Value cmpIm = SPIRVCompareOp::create(rewriter, loc, lhsIm, rhsIm);
152
153 rewriter.replaceOpWithNewOp<SPIRVCombinerOp>(op, spirvType, cmpRe, cmpIm);
154 return success();
155 }
156};
157
158struct MulOpPattern final : OpConversionPattern<complex::MulOp> {
159 using Base::Base;
160
161 LogicalResult
162 matchAndRewrite(complex::MulOp op, OpAdaptor adaptor,
163 ConversionPatternRewriter &rewriter) const override {
164 Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
165 if (!spirvType)
166 return rewriter.notifyMatchFailure(op, "unable to convert result type");
167
168 Location loc = op.getLoc();
169 Value lhs = adaptor.getLhs();
170 Value rhs = adaptor.getRhs();
171
172 Value a = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
173 Value b = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
174 Value c = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
175 Value d = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
176
177 Value ac = spirv::FMulOp::create(rewriter, loc, a, c);
178 Value bd = spirv::FMulOp::create(rewriter, loc, b, d);
179 Value ad = spirv::FMulOp::create(rewriter, loc, a, d);
180 Value bc = spirv::FMulOp::create(rewriter, loc, b, c);
181 Value resultRe = spirv::FSubOp::create(rewriter, loc, ac, bd);
182 Value resultIm = spirv::FAddOp::create(rewriter, loc, ad, bc);
183
184 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
185 op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
186 return success();
187 }
188};
189
190template <typename SqrtOp>
191struct AbsOpPattern final : OpConversionPattern<complex::AbsOp> {
192 using OpConversionPattern<complex::AbsOp>::OpConversionPattern;
193
194 LogicalResult
195 matchAndRewrite(complex::AbsOp op, OpAdaptor adaptor,
196 ConversionPatternRewriter &rewriter) const override {
197 Type spirvType =
198 this->getTypeConverter()->convertType(op.getResult().getType());
199 if (!spirvType)
200 return rewriter.notifyMatchFailure(op, "unable to convert result type");
201
202 Location loc = op.getLoc();
203 Value complexVal = adaptor.getComplex();
204
205 Value re =
206 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
207 Value im =
208 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
209
210 Value reSq = spirv::FMulOp::create(rewriter, loc, re, re);
211 Value imSq = spirv::FMulOp::create(rewriter, loc, im, im);
212 Value sum = spirv::FAddOp::create(rewriter, loc, reSq, imSq);
213
214 rewriter.replaceOpWithNewOp<SqrtOp>(op, sum);
215 return success();
216 }
217};
218
219template <typename ComplexOp, bool NegateReal>
220struct NegationOpPattern final : OpConversionPattern<ComplexOp> {
221 using OpConversionPattern<ComplexOp>::OpConversionPattern;
222 using OpAdaptor = typename ComplexOp::Adaptor;
223
224 LogicalResult
225 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
226 ConversionPatternRewriter &rewriter) const override {
227 Type spirvType =
228 this->getTypeConverter()->convertType(op.getResult().getType());
229 if (!spirvType)
230 return rewriter.notifyMatchFailure(op, "unable to convert result type");
231
232 Location loc = op.getLoc();
233 Value complexVal = adaptor.getComplex();
234
235 Value re =
236 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
237 Value im =
238 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
239
240 Value resultRe =
241 NegateReal ? spirv::FNegateOp::create(rewriter, loc, re) : re;
242 Value resultIm = spirv::FNegateOp::create(rewriter, loc, im);
243
244 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
245 op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
246 return success();
247 }
248};
249
250struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
251 using Base::Base;
252
253 LogicalResult
254 matchAndRewrite(complex::DivOp op, OpAdaptor adaptor,
255 ConversionPatternRewriter &rewriter) const override {
256 Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
257 if (!spirvType)
258 return rewriter.notifyMatchFailure(op, "unable to convert result type");
259
260 Location loc = op.getLoc();
261 Value lhs = adaptor.getLhs();
262 Value rhs = adaptor.getRhs();
263
264 Value a = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {0});
265 Value b = spirv::CompositeExtractOp::create(rewriter, loc, lhs, {1});
266 Value c = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {0});
267 Value d = spirv::CompositeExtractOp::create(rewriter, loc, rhs, {1});
268
269 Value ac = spirv::FMulOp::create(rewriter, loc, a, c);
270 Value bd = spirv::FMulOp::create(rewriter, loc, b, d);
271 Value bc = spirv::FMulOp::create(rewriter, loc, b, c);
272 Value ad = spirv::FMulOp::create(rewriter, loc, a, d);
273 Value cc = spirv::FMulOp::create(rewriter, loc, c, c);
274 Value dd = spirv::FMulOp::create(rewriter, loc, d, d);
275 Value denom = spirv::FAddOp::create(rewriter, loc, cc, dd);
276 Value numRe = spirv::FAddOp::create(rewriter, loc, ac, bd);
277 Value numIm = spirv::FSubOp::create(rewriter, loc, bc, ad);
278 Value resultRe = spirv::FDivOp::create(rewriter, loc, numRe, denom);
279 Value resultIm = spirv::FDivOp::create(rewriter, loc, numIm, denom);
280
281 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
282 op, spirvType, llvm::ArrayRef<Value>{resultRe, resultIm});
283 return success();
284 }
285};
286
287} // namespace
288
289//===----------------------------------------------------------------------===//
290// Pattern population
291//===----------------------------------------------------------------------===//
292
294 const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns) {
295 MLIRContext *context = patterns.getContext();
296
297 patterns.add<ConstantOpPattern, CreateOpPattern, ReOpPattern, ImOpPattern,
298 ElementwiseBinaryOpPattern<complex::AddOp, spirv::FAddOp>,
299 ElementwiseBinaryOpPattern<complex::SubOp, spirv::FSubOp>,
300 ComparisonOpPattern<complex::EqualOp, spirv::FOrdEqualOp,
301 spirv::LogicalAndOp>,
302 ComparisonOpPattern<complex::NotEqualOp, spirv::FUnordNotEqualOp,
303 spirv::LogicalOrOp>,
304 MulOpPattern, DivOpPattern,
305 NegationOpPattern<complex::NegOp, /*NegateReal=*/true>,
306 NegationOpPattern<complex::ConjOp, /*NegateReal=*/false>,
307 AbsOpPattern<spirv::GLSqrtOp>, AbsOpPattern<spirv::CLSqrtOp>>(
308 typeConverter, context);
309}
return success()
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Type conversion from builtin types to SPIR-V types for shader interface.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Include the generated interface declarations.
void populateComplexToSPIRVPatterns(const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns)
Appends to a pattern list additional patterns for translating Complex ops to SPIR-V ops.