19#define DEBUG_TYPE "complex-to-spirv-pattern"
29struct ConstantOpPattern final : OpConversionPattern<complex::ConstantOp> {
33 matchAndRewrite(complex::ConstantOp constOp, OpAdaptor adaptor,
34 ConversionPatternRewriter &rewriter)
const override {
36 getTypeConverter()->convertType<ShapedType>(constOp.getType());
38 return rewriter.notifyMatchFailure(constOp,
39 "unable to convert result type");
41 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
48struct CreateOpPattern final : OpConversionPattern<complex::CreateOp> {
52 matchAndRewrite(complex::CreateOp createOp, OpAdaptor adaptor,
53 ConversionPatternRewriter &rewriter)
const override {
54 Type spirvType = getTypeConverter()->convertType(createOp.getType());
56 return rewriter.notifyMatchFailure(createOp,
57 "unable to convert result type");
59 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
60 createOp, spirvType, adaptor.getOperands());
65struct ReOpPattern final : OpConversionPattern<complex::ReOp> {
69 matchAndRewrite(complex::ReOp reOp, OpAdaptor adaptor,
70 ConversionPatternRewriter &rewriter)
const override {
71 Type spirvType = getTypeConverter()->convertType(reOp.getType());
73 return rewriter.notifyMatchFailure(reOp,
"unable to convert result type");
75 rewriter.replaceOpWithNewOp<spirv::CompositeExtractOp>(
81struct ImOpPattern final : OpConversionPattern<complex::ImOp> {
85 matchAndRewrite(complex::ImOp imOp, OpAdaptor adaptor,
86 ConversionPatternRewriter &rewriter)
const override {
87 Type spirvType = getTypeConverter()->convertType(imOp.getType());
89 return rewriter.notifyMatchFailure(imOp,
"unable to convert result type");
91 rewriter.replaceOpWithNewOp<spirv::CompositeExtractOp>(
97template <
typename ComplexOp,
typename SPIRVOp>
98struct ElementwiseBinaryOpPattern final : OpConversionPattern<ComplexOp> {
99 using OpConversionPattern<ComplexOp>::OpConversionPattern;
100 using OpAdaptor =
typename ComplexOp::Adaptor;
103 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
104 ConversionPatternRewriter &rewriter)
const override {
106 this->getTypeConverter()->convertType(op.getResult().getType());
108 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
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});
119 Value resultRe = SPIRVOp::create(rewriter, loc, lhsRe, rhsRe);
120 Value resultIm = SPIRVOp::create(rewriter, loc, lhsIm, rhsIm);
122 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
128template <
typename ComplexOp,
typename SPIRVCompareOp,
typename SPIRVCombinerOp>
129struct ComparisonOpPattern final : OpConversionPattern<ComplexOp> {
130 using OpConversionPattern<ComplexOp>::OpConversionPattern;
131 using OpAdaptor =
typename ComplexOp::Adaptor;
134 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
135 ConversionPatternRewriter &rewriter)
const override {
137 this->getTypeConverter()->convertType(op.getResult().getType());
139 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
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});
150 Value cmpRe = SPIRVCompareOp::create(rewriter, loc, lhsRe, rhsRe);
151 Value cmpIm = SPIRVCompareOp::create(rewriter, loc, lhsIm, rhsIm);
153 rewriter.replaceOpWithNewOp<SPIRVCombinerOp>(op, spirvType, cmpRe, cmpIm);
158struct MulOpPattern final : OpConversionPattern<complex::MulOp> {
162 matchAndRewrite(complex::MulOp op, OpAdaptor adaptor,
163 ConversionPatternRewriter &rewriter)
const override {
164 Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
166 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
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});
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);
184 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
190template <
typename SqrtOp>
191struct AbsOpPattern final : OpConversionPattern<complex::AbsOp> {
192 using OpConversionPattern<complex::AbsOp>::OpConversionPattern;
195 matchAndRewrite(complex::AbsOp op, OpAdaptor adaptor,
196 ConversionPatternRewriter &rewriter)
const override {
198 this->getTypeConverter()->convertType(op.getResult().getType());
200 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
203 Value complexVal = adaptor.getComplex();
206 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
208 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
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);
214 rewriter.replaceOpWithNewOp<SqrtOp>(op, sum);
219template <
typename ComplexOp,
bool NegateReal>
220struct NegationOpPattern final : OpConversionPattern<ComplexOp> {
221 using OpConversionPattern<ComplexOp>::OpConversionPattern;
222 using OpAdaptor =
typename ComplexOp::Adaptor;
225 matchAndRewrite(ComplexOp op, OpAdaptor adaptor,
226 ConversionPatternRewriter &rewriter)
const override {
228 this->getTypeConverter()->convertType(op.getResult().getType());
230 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
233 Value complexVal = adaptor.getComplex();
236 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {0});
238 spirv::CompositeExtractOp::create(rewriter, loc, complexVal, {1});
241 NegateReal ? spirv::FNegateOp::create(rewriter, loc, re) : re;
242 Value resultIm = spirv::FNegateOp::create(rewriter, loc, im);
244 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
250struct DivOpPattern final : OpConversionPattern<complex::DivOp> {
254 matchAndRewrite(complex::DivOp op, OpAdaptor adaptor,
255 ConversionPatternRewriter &rewriter)
const override {
256 Type spirvType = getTypeConverter()->convertType(op.getResult().getType());
258 return rewriter.notifyMatchFailure(op,
"unable to convert result type");
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});
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);
281 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(
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,
304 MulOpPattern, DivOpPattern,
305 NegationOpPattern<complex::NegOp,
true>,
306 NegationOpPattern<complex::ConjOp,
false>,
307 AbsOpPattern<spirv::GLSqrtOp>, AbsOpPattern<spirv::CLSqrtOp>>(
308 typeConverter, context);
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...
MLIRContext is the top-level object for a collection of MLIR operations.
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
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.