MLIR 24.0.0git
ComplexToStandard.cpp
Go to the documentation of this file.
1//===- ComplexToStandard.cpp - conversion from Complex to Standard dialect ===//
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
10
17#include <type_traits>
18
19namespace mlir {
20#define GEN_PASS_DEF_CONVERTCOMPLEXTOSTANDARDPASS
21#include "mlir/Conversion/Passes.h.inc"
22} // namespace mlir
23
24using namespace mlir;
25
26namespace {
27
28enum class AbsFn { abs, sqrt, rsqrt };
29
30// Returns the absolute value, its square root or its reciprocal square root.
31Value computeAbs(Value real, Value imag, arith::FastMathFlags fmf,
32 ImplicitLocOpBuilder &b, AbsFn fn = AbsFn::abs) {
33 Value one = arith::ConstantOp::create(b, real.getType(),
34 b.getFloatAttr(real.getType(), 1.0));
35
36 Value absReal = math::AbsFOp::create(b, real, fmf);
37 Value absImag = math::AbsFOp::create(b, imag, fmf);
38
39 Value max = arith::MaximumFOp::create(b, absReal, absImag, fmf);
40 Value min = arith::MinimumFOp::create(b, absReal, absImag, fmf);
41
42 // The lowering below requires NaNs and infinities to work correctly.
43 arith::FastMathFlags fmfWithNaNInf = arith::bitEnumClear(
44 fmf, arith::FastMathFlags::nnan | arith::FastMathFlags::ninf);
45 Value ratio = arith::DivFOp::create(b, min, max, fmfWithNaNInf);
46 Value ratioSq = arith::MulFOp::create(b, ratio, ratio, fmfWithNaNInf);
47 Value ratioSqPlusOne = arith::AddFOp::create(b, ratioSq, one, fmfWithNaNInf);
49
50 if (fn == AbsFn::rsqrt) {
51 ratioSqPlusOne = math::RsqrtOp::create(b, ratioSqPlusOne, fmfWithNaNInf);
52 min = math::RsqrtOp::create(b, min, fmfWithNaNInf);
53 max = math::RsqrtOp::create(b, max, fmfWithNaNInf);
54 }
55
56 if (fn == AbsFn::sqrt) {
57 Value quarter = arith::ConstantOp::create(
58 b, real.getType(), b.getFloatAttr(real.getType(), 0.25));
59 // sqrt(sqrt(a*b)) would avoid the pow, but will overflow more easily.
60 Value sqrt = math::SqrtOp::create(b, max, fmfWithNaNInf);
61 Value p025 =
62 math::PowFOp::create(b, ratioSqPlusOne, quarter, fmfWithNaNInf);
63 result = arith::MulFOp::create(b, sqrt, p025, fmfWithNaNInf);
64 } else {
65 Value sqrt = math::SqrtOp::create(b, ratioSqPlusOne, fmfWithNaNInf);
66 result = arith::MulFOp::create(b, max, sqrt, fmfWithNaNInf);
67 }
68
69 Value isNaN = arith::CmpFOp::create(b, arith::CmpFPredicate::UNO, result,
70 result, fmfWithNaNInf);
71 return arith::SelectOp::create(b, isNaN, min, result);
72}
73
74struct AbsOpConversion : public OpConversionPattern<complex::AbsOp> {
75 using OpConversionPattern<complex::AbsOp>::OpConversionPattern;
76
77 LogicalResult
78 matchAndRewrite(complex::AbsOp op, OpAdaptor adaptor,
79 ConversionPatternRewriter &rewriter) const override {
80 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
81
82 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
83
84 Value real = complex::ReOp::create(b, adaptor.getComplex());
85 Value imag = complex::ImOp::create(b, adaptor.getComplex());
86 rewriter.replaceOp(op, computeAbs(real, imag, fmf, b));
87
88 return success();
89 }
90};
91
92// atan2(y,x) = -i * log((x + i * y)/sqrt(x**2+y**2))
93struct Atan2OpConversion : public OpConversionPattern<complex::Atan2Op> {
94 using OpConversionPattern<complex::Atan2Op>::OpConversionPattern;
95
96 LogicalResult
97 matchAndRewrite(complex::Atan2Op op, OpAdaptor adaptor,
98 ConversionPatternRewriter &rewriter) const override {
99 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
100
101 auto type = cast<ComplexType>(op.getType());
102 Type elementType = type.getElementType();
103 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
104
105 Value lhs = adaptor.getLhs();
106 Value rhs = adaptor.getRhs();
107
108 Value rhsSquared = complex::MulOp::create(b, type, rhs, rhs, fmf);
109 Value lhsSquared = complex::MulOp::create(b, type, lhs, lhs, fmf);
110 Value rhsSquaredPlusLhsSquared =
111 complex::AddOp::create(b, type, rhsSquared, lhsSquared, fmf);
112 Value sqrtOfRhsSquaredPlusLhsSquared =
113 complex::SqrtOp::create(b, type, rhsSquaredPlusLhsSquared, fmf);
114
115 Value zero =
116 arith::ConstantOp::create(b, elementType, b.getZeroAttr(elementType));
117 Value one = arith::ConstantOp::create(b, elementType,
118 b.getFloatAttr(elementType, 1));
119 Value i = complex::CreateOp::create(b, type, zero, one);
120 Value iTimesLhs = complex::MulOp::create(b, i, lhs, fmf);
121 Value rhsPlusILhs = complex::AddOp::create(b, rhs, iTimesLhs, fmf);
122
123 Value divResult = complex::DivOp::create(
124 b, rhsPlusILhs, sqrtOfRhsSquaredPlusLhsSquared, fmf);
125 Value logResult = complex::LogOp::create(b, divResult, fmf);
126
127 Value negativeOne = arith::ConstantOp::create(
128 b, elementType, b.getFloatAttr(elementType, -1));
129 Value negativeI = complex::CreateOp::create(b, type, zero, negativeOne);
130
131 rewriter.replaceOpWithNewOp<complex::MulOp>(op, negativeI, logResult, fmf);
132 return success();
133 }
134};
135
136template <typename ComparisonOp, arith::CmpFPredicate p>
137struct ComparisonOpConversion : public OpConversionPattern<ComparisonOp> {
138 using OpConversionPattern<ComparisonOp>::OpConversionPattern;
139 using ResultCombiner =
140 std::conditional_t<std::is_same<ComparisonOp, complex::EqualOp>::value,
141 arith::AndIOp, arith::OrIOp>;
142
143 LogicalResult
144 matchAndRewrite(ComparisonOp op, typename ComparisonOp::Adaptor adaptor,
145 ConversionPatternRewriter &rewriter) const override {
146 auto loc = op.getLoc();
147 auto type = cast<ComplexType>(adaptor.getLhs().getType()).getElementType();
148
149 Value realLhs =
150 complex::ReOp::create(rewriter, loc, type, adaptor.getLhs());
151 Value imagLhs =
152 complex::ImOp::create(rewriter, loc, type, adaptor.getLhs());
153 Value realRhs =
154 complex::ReOp::create(rewriter, loc, type, adaptor.getRhs());
155 Value imagRhs =
156 complex::ImOp::create(rewriter, loc, type, adaptor.getRhs());
157 Value realComparison =
158 arith::CmpFOp::create(rewriter, loc, p, realLhs, realRhs);
159 Value imagComparison =
160 arith::CmpFOp::create(rewriter, loc, p, imagLhs, imagRhs);
161
162 rewriter.replaceOpWithNewOp<ResultCombiner>(op, realComparison,
163 imagComparison);
164 return success();
165 }
166};
167
168// Default conversion which applies the BinaryStandardOp separately on the real
169// and imaginary parts. Can for example be used for complex::AddOp and
170// complex::SubOp.
171template <typename BinaryComplexOp, typename BinaryStandardOp>
172struct BinaryComplexOpConversion : public OpConversionPattern<BinaryComplexOp> {
173 using OpConversionPattern<BinaryComplexOp>::OpConversionPattern;
174
175 LogicalResult
176 matchAndRewrite(BinaryComplexOp op, typename BinaryComplexOp::Adaptor adaptor,
177 ConversionPatternRewriter &rewriter) const override {
178 auto type = cast<ComplexType>(adaptor.getLhs().getType());
179 auto elementType = cast<FloatType>(type.getElementType());
180 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
181 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
182
183 Value realLhs = complex::ReOp::create(b, elementType, adaptor.getLhs());
184 Value realRhs = complex::ReOp::create(b, elementType, adaptor.getRhs());
185 Value resultReal = BinaryStandardOp::create(b, elementType, realLhs,
186 realRhs, fmf.getValue());
187 Value imagLhs = complex::ImOp::create(b, elementType, adaptor.getLhs());
188 Value imagRhs = complex::ImOp::create(b, elementType, adaptor.getRhs());
189 Value resultImag = BinaryStandardOp::create(b, elementType, imagLhs,
190 imagRhs, fmf.getValue());
191 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
192 resultImag);
193 return success();
194 }
195};
196
197template <typename TrigonometricOp>
198struct TrigonometricOpConversion : public OpConversionPattern<TrigonometricOp> {
199 using OpAdaptor = typename OpConversionPattern<TrigonometricOp>::OpAdaptor;
200
201 using OpConversionPattern<TrigonometricOp>::OpConversionPattern;
202
203 LogicalResult
204 matchAndRewrite(TrigonometricOp op, OpAdaptor adaptor,
205 ConversionPatternRewriter &rewriter) const override {
206 auto loc = op.getLoc();
207 auto type = cast<ComplexType>(adaptor.getComplex().getType());
208 auto elementType = cast<FloatType>(type.getElementType());
209 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
210
211 Value real =
212 complex::ReOp::create(rewriter, loc, elementType, adaptor.getComplex());
213 Value imag =
214 complex::ImOp::create(rewriter, loc, elementType, adaptor.getComplex());
215
216 // Trigonometric ops use a set of common building blocks to convert to real
217 // ops. Here we create these building blocks and call into an op-specific
218 // implementation in the subclass to combine them.
219 Value half = arith::ConstantOp::create(
220 rewriter, loc, elementType, rewriter.getFloatAttr(elementType, 0.5));
221 Value exp = math::ExpOp::create(rewriter, loc, imag, fmf);
222 Value scaledExp = arith::MulFOp::create(rewriter, loc, half, exp, fmf);
223 Value reciprocalExp = arith::DivFOp::create(rewriter, loc, half, exp, fmf);
224 Value sin = math::SinOp::create(rewriter, loc, real, fmf);
225 Value cos = math::CosOp::create(rewriter, loc, real, fmf);
226
227 auto resultPair =
228 combine(loc, scaledExp, reciprocalExp, sin, cos, rewriter, fmf);
229
230 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultPair.first,
231 resultPair.second);
232 return success();
233 }
234
235 virtual std::pair<Value, Value>
236 combine(Location loc, Value scaledExp, Value reciprocalExp, Value sin,
237 Value cos, ConversionPatternRewriter &rewriter,
238 arith::FastMathFlagsAttr fmf) const = 0;
239};
240
241struct CosOpConversion : public TrigonometricOpConversion<complex::CosOp> {
242 using TrigonometricOpConversion<complex::CosOp>::TrigonometricOpConversion;
243
244 std::pair<Value, Value> combine(Location loc, Value scaledExp,
245 Value reciprocalExp, Value sin, Value cos,
246 ConversionPatternRewriter &rewriter,
247 arith::FastMathFlagsAttr fmf) const override {
248 // Complex cosine is defined as;
249 // cos(x + iy) = 0.5 * (exp(i(x + iy)) + exp(-i(x + iy)))
250 // Plugging in:
251 // exp(i(x+iy)) = exp(-y + ix) = exp(-y)(cos(x) + i sin(x))
252 // exp(-i(x+iy)) = exp(y + i(-x)) = exp(y)(cos(x) + i (-sin(x)))
253 // and defining t := exp(y)
254 // We get:
255 // Re(cos(x + iy)) = (0.5/t + 0.5*t) * cos x
256 // Im(cos(x + iy)) = (0.5/t - 0.5*t) * sin x
257 Value sum =
258 arith::AddFOp::create(rewriter, loc, reciprocalExp, scaledExp, fmf);
259 Value resultReal = arith::MulFOp::create(rewriter, loc, sum, cos, fmf);
260 Value diff =
261 arith::SubFOp::create(rewriter, loc, reciprocalExp, scaledExp, fmf);
262 Value resultImag = arith::MulFOp::create(rewriter, loc, diff, sin, fmf);
263 return {resultReal, resultImag};
264 }
265};
266
267struct DivOpConversion : public OpConversionPattern<complex::DivOp> {
268 DivOpConversion(MLIRContext *context, complex::ComplexRangeFlags target)
269 : OpConversionPattern<complex::DivOp>(context), complexRange(target) {}
270
271 using OpConversionPattern<complex::DivOp>::OpConversionPattern;
272
273 LogicalResult
274 matchAndRewrite(complex::DivOp op, OpAdaptor adaptor,
275 ConversionPatternRewriter &rewriter) const override {
276 auto loc = op.getLoc();
277 auto type = cast<ComplexType>(adaptor.getLhs().getType());
278 auto elementType = cast<FloatType>(type.getElementType());
279 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
280
281 Value lhsReal =
282 complex::ReOp::create(rewriter, loc, elementType, adaptor.getLhs());
283 Value lhsImag =
284 complex::ImOp::create(rewriter, loc, elementType, adaptor.getLhs());
285 Value rhsReal =
286 complex::ReOp::create(rewriter, loc, elementType, adaptor.getRhs());
287 Value rhsImag =
288 complex::ImOp::create(rewriter, loc, elementType, adaptor.getRhs());
289
290 Value resultReal, resultImag;
291
292 if (complexRange == complex::ComplexRangeFlags::basic ||
293 complexRange == complex::ComplexRangeFlags::none) {
295 rewriter, loc, lhsReal, lhsImag, rhsReal, rhsImag, fmf, &resultReal,
296 &resultImag);
297 } else if (complexRange == complex::ComplexRangeFlags::improved) {
299 rewriter, loc, lhsReal, lhsImag, rhsReal, rhsImag, fmf, &resultReal,
300 &resultImag);
301 }
302
303 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
304 resultImag);
305
306 return success();
307 }
308
309private:
310 complex::ComplexRangeFlags complexRange;
311};
312
313struct ExpOpConversion : public OpConversionPattern<complex::ExpOp> {
314 using OpConversionPattern<complex::ExpOp>::OpConversionPattern;
315
316 // exp(x+I*y) = exp(x)*(cos(y)+I*sin(y))
317 // Handle special cases as StableHLO implementation does:
318 // 1. When b == 0, set imag(exp(z)) = 0
319 // 2. When exp(x) == inf, use exp(x/2)*(cos(y)+I*sin(y))*exp(x/2)
320 LogicalResult
321 matchAndRewrite(complex::ExpOp op, OpAdaptor adaptor,
322 ConversionPatternRewriter &rewriter) const override {
323 auto loc = op.getLoc();
324 auto type = cast<ComplexType>(adaptor.getComplex().getType());
325 auto ET = cast<FloatType>(type.getElementType());
326 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
327 const auto &floatSemantics = ET.getFloatSemantics();
328 ImplicitLocOpBuilder b(loc, rewriter);
329
330 Value x = complex::ReOp::create(b, ET, adaptor.getComplex());
331 Value y = complex::ImOp::create(b, ET, adaptor.getComplex());
332 Value zero = arith::ConstantOp::create(b, ET, b.getZeroAttr(ET));
333 Value half = arith::ConstantOp::create(b, ET, b.getFloatAttr(ET, 0.5));
334 Value inf = arith::ConstantOp::create(
335 b, ET, b.getFloatAttr(ET, APFloat::getInf(floatSemantics)));
336
337 Value exp = math::ExpOp::create(b, x, fmf);
338 Value xHalf = arith::MulFOp::create(b, x, half, fmf);
339 Value expHalf = math::ExpOp::create(b, xHalf, fmf);
340 Value cos = math::CosOp::create(b, y, fmf);
341 Value sin = math::SinOp::create(b, y, fmf);
342
343 Value expIsInf =
344 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, exp, inf, fmf);
345 Value yIsZero =
346 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, y, zero);
347
348 // Real path: select between exp(x)*cos(y) and exp(x/2)*cos(y)*exp(x/2)
349 Value realNormal = arith::MulFOp::create(b, exp, cos, fmf);
350 Value expHalfCos = arith::MulFOp::create(b, expHalf, cos, fmf);
351 Value realOverflow = arith::MulFOp::create(b, expHalfCos, expHalf, fmf);
352 Value resultReal =
353 arith::SelectOp::create(b, expIsInf, realOverflow, realNormal);
354
355 // Imaginary part: if y == 0 return 0 else select between exp(x)*sin(y) and
356 // exp(x/2)*sin(y)*exp(x/2)
357 Value imagNormal = arith::MulFOp::create(b, exp, sin, fmf);
358 Value expHalfSin = arith::MulFOp::create(b, expHalf, sin, fmf);
359 Value imagOverflow = arith::MulFOp::create(b, expHalfSin, expHalf, fmf);
360 Value imagNonZero =
361 arith::SelectOp::create(b, expIsInf, imagOverflow, imagNormal);
362 Value resultImag = arith::SelectOp::create(b, yIsZero, zero, imagNonZero);
363
364 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
365 resultImag);
366 return success();
367 }
368};
369
370Value evaluatePolynomial(ImplicitLocOpBuilder &b, Value arg,
371 ArrayRef<double> coefficients,
372 arith::FastMathFlagsAttr fmf) {
373 auto argType = mlir::cast<FloatType>(arg.getType());
374 Value poly =
375 arith::ConstantOp::create(b, b.getFloatAttr(argType, coefficients[0]));
376 for (unsigned i = 1; i < coefficients.size(); ++i) {
377 poly = math::FmaOp::create(
378 b, poly, arg,
379 arith::ConstantOp::create(b, b.getFloatAttr(argType, coefficients[i])),
380 fmf);
381 }
382 return poly;
383}
384
385struct Expm1OpConversion : public OpConversionPattern<complex::Expm1Op> {
386 using OpConversionPattern<complex::Expm1Op>::OpConversionPattern;
387
388 // e^(a+bi)-1 = (e^a*cos(b)-1)+e^a*sin(b)i
389 // [handle inaccuracies when a and/or b are small]
390 // = ((e^a - 1) * cos(b) + cos(b) - 1) + e^a*sin(b)i
391 // = (expm1(a) * cos(b) + cosm1(b)) + e^a*sin(b)i
392 LogicalResult
393 matchAndRewrite(complex::Expm1Op op, OpAdaptor adaptor,
394 ConversionPatternRewriter &rewriter) const override {
395 auto type = op.getType();
396 auto elemType = mlir::cast<FloatType>(type.getElementType());
397
398 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
399 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
400 Value real = complex::ReOp::create(b, adaptor.getComplex());
401 Value imag = complex::ImOp::create(b, adaptor.getComplex());
402
403 Value zero = arith::ConstantOp::create(b, b.getFloatAttr(elemType, 0.0));
404 Value one = arith::ConstantOp::create(b, b.getFloatAttr(elemType, 1.0));
405
406 Value expm1Real = math::ExpM1Op::create(b, real, fmf);
407 Value expReal = arith::AddFOp::create(b, expm1Real, one, fmf);
408
409 Value sinImag = math::SinOp::create(b, imag, fmf);
410 Value cosm1Imag = emitCosm1(imag, fmf, b);
411 Value cosImag = arith::AddFOp::create(b, cosm1Imag, one, fmf);
412
413 Value realResult = arith::AddFOp::create(
414 b, arith::MulFOp::create(b, expm1Real, cosImag, fmf), cosm1Imag, fmf);
415
416 Value imagIsZero = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, imag,
417 zero, fmf.getValue());
418 Value imagResult = arith::SelectOp::create(
419 b, imagIsZero, zero, arith::MulFOp::create(b, expReal, sinImag, fmf));
420
421 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, realResult,
422 imagResult);
423 return success();
424 }
425
426private:
427 Value emitCosm1(Value arg, arith::FastMathFlagsAttr fmf,
428 ImplicitLocOpBuilder &b) const {
429 auto argType = mlir::cast<FloatType>(arg.getType());
430 auto negHalf = arith::ConstantOp::create(b, b.getFloatAttr(argType, -0.5));
431 auto one = arith::ConstantOp::create(b, b.getFloatAttr(argType, 1.0));
432
433 // Algorithm copied from cephes cosm1.
434 SmallVector<double, 7> kCoeffs{
435 4.7377507964246204691685E-14, -1.1470284843425359765671E-11,
436 2.0876754287081521758361E-9, -2.7557319214999787979814E-7,
437 2.4801587301570552304991E-5, -1.3888888888888872993737E-3,
438 4.1666666666666666609054E-2,
439 };
440 Value cos = math::CosOp::create(b, arg, fmf);
441 Value forLargeArg = arith::SubFOp::create(b, cos, one, fmf);
442
443 Value argPow2 = arith::MulFOp::create(b, arg, arg, fmf);
444 Value argPow4 = arith::MulFOp::create(b, argPow2, argPow2, fmf);
445 Value poly = evaluatePolynomial(b, argPow2, kCoeffs, fmf);
446
447 auto forSmallArg =
448 arith::AddFOp::create(b, arith::MulFOp::create(b, argPow4, poly, fmf),
449 arith::MulFOp::create(b, negHalf, argPow2, fmf));
450
451 // (pi/4)^2 is approximately 0.61685
452 Value piOver4Pow2 =
453 arith::ConstantOp::create(b, b.getFloatAttr(argType, 0.61685));
454 Value cond = arith::CmpFOp::create(b, arith::CmpFPredicate::OGE, argPow2,
455 piOver4Pow2, fmf.getValue());
456 return arith::SelectOp::create(b, cond, forLargeArg, forSmallArg);
457 }
458};
459
460struct LogOpConversion : public OpConversionPattern<complex::LogOp> {
461 using OpConversionPattern<complex::LogOp>::OpConversionPattern;
462
463 LogicalResult
464 matchAndRewrite(complex::LogOp op, OpAdaptor adaptor,
465 ConversionPatternRewriter &rewriter) const override {
466 auto type = cast<ComplexType>(adaptor.getComplex().getType());
467 auto elementType = cast<FloatType>(type.getElementType());
468 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
469 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
470
471 Value abs = complex::AbsOp::create(b, elementType, adaptor.getComplex(),
472 fmf.getValue());
473 Value resultReal = math::LogOp::create(b, elementType, abs, fmf.getValue());
474 Value real = complex::ReOp::create(b, elementType, adaptor.getComplex());
475 Value imag = complex::ImOp::create(b, elementType, adaptor.getComplex());
476 Value resultImag =
477 math::Atan2Op::create(b, elementType, imag, real, fmf.getValue());
478 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
479 resultImag);
480 return success();
481 }
482};
483
484struct Log1pOpConversion : public OpConversionPattern<complex::Log1pOp> {
485 using OpConversionPattern<complex::Log1pOp>::OpConversionPattern;
486
487 LogicalResult
488 matchAndRewrite(complex::Log1pOp op, OpAdaptor adaptor,
489 ConversionPatternRewriter &rewriter) const override {
490 auto type = cast<ComplexType>(adaptor.getComplex().getType());
491 auto elementType = cast<FloatType>(type.getElementType());
492 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
493 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
494
495 Value real = complex::ReOp::create(b, adaptor.getComplex());
496 Value imag = complex::ImOp::create(b, adaptor.getComplex());
497
498 Value half = arith::ConstantOp::create(b, elementType,
499 b.getFloatAttr(elementType, 0.5));
500 Value one = arith::ConstantOp::create(b, elementType,
501 b.getFloatAttr(elementType, 1));
502 Value realPlusOne = arith::AddFOp::create(b, real, one, fmf);
503 Value absRealPlusOne = math::AbsFOp::create(b, realPlusOne, fmf);
504 Value absImag = math::AbsFOp::create(b, imag, fmf);
505
506 Value maxAbs = arith::MaximumFOp::create(b, absRealPlusOne, absImag, fmf);
507 Value minAbs = arith::MinimumFOp::create(b, absRealPlusOne, absImag, fmf);
508
509 Value useReal = arith::CmpFOp::create(b, arith::CmpFPredicate::OGT,
510 realPlusOne, absImag, fmf);
511 Value maxMinusOne = arith::SubFOp::create(b, maxAbs, one, fmf);
512 Value maxAbsOfRealPlusOneAndImagMinusOne =
513 arith::SelectOp::create(b, useReal, real, maxMinusOne);
514 arith::FastMathFlags fmfWithNaNInf = arith::bitEnumClear(
515 fmf, arith::FastMathFlags::nnan | arith::FastMathFlags::ninf);
516 Value minMaxRatio = arith::DivFOp::create(b, minAbs, maxAbs, fmfWithNaNInf);
517 Value logOfMaxAbsOfRealPlusOneAndImag =
518 math::Log1pOp::create(b, maxAbsOfRealPlusOneAndImagMinusOne, fmf);
519 Value logOfSqrtPart = math::Log1pOp::create(
520 b, arith::MulFOp::create(b, minMaxRatio, minMaxRatio, fmfWithNaNInf),
521 fmfWithNaNInf);
522 Value r = arith::AddFOp::create(
523 b, arith::MulFOp::create(b, half, logOfSqrtPart, fmfWithNaNInf),
524 logOfMaxAbsOfRealPlusOneAndImag, fmfWithNaNInf);
525 Value resultReal = arith::SelectOp::create(
526 b,
527 arith::CmpFOp::create(b, arith::CmpFPredicate::UNO, r, r,
528 fmfWithNaNInf),
529 minAbs, r);
530 Value resultImag = math::Atan2Op::create(b, imag, realPlusOne, fmf);
531 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
532 resultImag);
533 return success();
534 }
535};
536
537struct MulOpConversion : public OpConversionPattern<complex::MulOp> {
538 using OpConversionPattern<complex::MulOp>::OpConversionPattern;
539
540 LogicalResult
541 matchAndRewrite(complex::MulOp op, OpAdaptor adaptor,
542 ConversionPatternRewriter &rewriter) const override {
543 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
544 auto type = cast<ComplexType>(adaptor.getLhs().getType());
545 auto elementType = cast<FloatType>(type.getElementType());
546 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
547 auto fmfValue = fmf.getValue();
548 Value lhsReal = complex::ReOp::create(b, elementType, adaptor.getLhs());
549 Value lhsImag = complex::ImOp::create(b, elementType, adaptor.getLhs());
550 Value rhsReal = complex::ReOp::create(b, elementType, adaptor.getRhs());
551 Value rhsImag = complex::ImOp::create(b, elementType, adaptor.getRhs());
552 Value real;
553 Value imag;
554 if (arith::bitEnumContainsAll(fmfValue, arith::FastMathFlags::contract)) {
555 Value lhsImagTimesRhsImag =
556 arith::MulFOp::create(b, lhsImag, rhsImag, fmfValue);
557 Value negLhsImagTimesRhsImag =
558 arith::NegFOp::create(b, lhsImagTimesRhsImag, fmfValue);
559 real = math::FmaOp::create(b, lhsReal, rhsReal, negLhsImagTimesRhsImag,
560 fmfValue);
561
562 Value lhsImagTimesRhsReal =
563 arith::MulFOp::create(b, lhsImag, rhsReal, fmfValue);
564 imag = math::FmaOp::create(b, lhsReal, rhsImag, lhsImagTimesRhsReal,
565 fmfValue);
566 } else {
567 Value lhsRealTimesRhsReal =
568 arith::MulFOp::create(b, lhsReal, rhsReal, fmfValue);
569 Value lhsImagTimesRhsImag =
570 arith::MulFOp::create(b, lhsImag, rhsImag, fmfValue);
571 Value lhsImagTimesRhsReal =
572 arith::MulFOp::create(b, lhsImag, rhsReal, fmfValue);
573 Value lhsRealTimesRhsImag =
574 arith::MulFOp::create(b, lhsReal, rhsImag, fmfValue);
575
576 real = arith::SubFOp::create(b, lhsRealTimesRhsReal, lhsImagTimesRhsImag,
577 fmfValue);
578 imag = arith::AddFOp::create(b, lhsImagTimesRhsReal, lhsRealTimesRhsImag,
579 fmfValue);
580 }
581 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, real, imag);
582 return success();
583 }
584};
585
586struct NegOpConversion : public OpConversionPattern<complex::NegOp> {
587 using OpConversionPattern<complex::NegOp>::OpConversionPattern;
588
589 LogicalResult
590 matchAndRewrite(complex::NegOp op, OpAdaptor adaptor,
591 ConversionPatternRewriter &rewriter) const override {
592 auto loc = op.getLoc();
593 auto type = cast<ComplexType>(adaptor.getComplex().getType());
594 auto elementType = cast<FloatType>(type.getElementType());
595 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
596
597 Value real =
598 complex::ReOp::create(rewriter, loc, elementType, adaptor.getComplex());
599 Value imag =
600 complex::ImOp::create(rewriter, loc, elementType, adaptor.getComplex());
601 Value negReal = arith::NegFOp::create(rewriter, loc, real, fmf);
602 Value negImag = arith::NegFOp::create(rewriter, loc, imag, fmf);
603 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, negReal, negImag);
604 return success();
605 }
606};
607
608struct SinOpConversion : public TrigonometricOpConversion<complex::SinOp> {
609 using TrigonometricOpConversion<complex::SinOp>::TrigonometricOpConversion;
610
611 std::pair<Value, Value> combine(Location loc, Value scaledExp,
612 Value reciprocalExp, Value sin, Value cos,
613 ConversionPatternRewriter &rewriter,
614 arith::FastMathFlagsAttr fmf) const override {
615 // Complex sine is defined as;
616 // sin(x + iy) = -0.5i * (exp(i(x + iy)) - exp(-i(x + iy)))
617 // Plugging in:
618 // exp(i(x+iy)) = exp(-y + ix) = exp(-y)(cos(x) + i sin(x))
619 // exp(-i(x+iy)) = exp(y + i(-x)) = exp(y)(cos(x) + i (-sin(x)))
620 // and defining t := exp(y)
621 // We get:
622 // Re(sin(x + iy)) = (0.5*t + 0.5/t) * sin x
623 // Im(sin(x + iy)) = (0.5*t - 0.5/t) * cos x
624 Value sum =
625 arith::AddFOp::create(rewriter, loc, scaledExp, reciprocalExp, fmf);
626 Value resultReal = arith::MulFOp::create(rewriter, loc, sum, sin, fmf);
627 Value diff =
628 arith::SubFOp::create(rewriter, loc, scaledExp, reciprocalExp, fmf);
629 Value resultImag = arith::MulFOp::create(rewriter, loc, diff, cos, fmf);
630 return {resultReal, resultImag};
631 }
632};
633
634// The algorithm is listed in https://dl.acm.org/doi/pdf/10.1145/363717.363780.
635struct SqrtOpConversion : public OpConversionPattern<complex::SqrtOp> {
636 using OpConversionPattern<complex::SqrtOp>::OpConversionPattern;
637
638 LogicalResult
639 matchAndRewrite(complex::SqrtOp op, OpAdaptor adaptor,
640 ConversionPatternRewriter &rewriter) const override {
641 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
642
643 auto type = cast<ComplexType>(op.getType());
644 auto elementType = cast<FloatType>(type.getElementType());
645 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
646
647 auto cst = [&](APFloat v) {
648 return arith::ConstantOp::create(b, elementType,
649 b.getFloatAttr(elementType, v));
650 };
651 const auto &floatSemantics = elementType.getFloatSemantics();
652 Value zero = cst(APFloat::getZero(floatSemantics));
653 Value half = arith::ConstantOp::create(b, elementType,
654 b.getFloatAttr(elementType, 0.5));
655
656 Value real = complex::ReOp::create(b, elementType, adaptor.getComplex());
657 Value imag = complex::ImOp::create(b, elementType, adaptor.getComplex());
658 Value absSqrt = computeAbs(real, imag, fmf, b, AbsFn::sqrt);
659 Value argArg = math::Atan2Op::create(b, imag, real, fmf);
660 Value sqrtArg = arith::MulFOp::create(b, argArg, half, fmf);
661 Value cos = math::CosOp::create(b, sqrtArg, fmf);
662 Value sin = math::SinOp::create(b, sqrtArg, fmf);
663 // sin(atan2(0, inf)) = 0, sqrt(abs(inf)) = inf, but we can't multiply
664 // 0 * inf.
665 Value sinIsZero =
666 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, sin, zero, fmf);
667
668 Value resultReal = arith::MulFOp::create(b, absSqrt, cos, fmf);
669 Value resultImag = arith::SelectOp::create(
670 b, sinIsZero, zero, arith::MulFOp::create(b, absSqrt, sin, fmf));
671 if (!arith::bitEnumContainsAll(fmf, arith::FastMathFlags::nnan |
672 arith::FastMathFlags::ninf)) {
673 Value inf = cst(APFloat::getInf(floatSemantics));
674 Value negInf = cst(APFloat::getInf(floatSemantics, true));
675 Value nan = cst(APFloat::getNaN(floatSemantics));
676 Value absImag = math::AbsFOp::create(b, elementType, imag, fmf);
677
678 Value absImagIsInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
679 absImag, inf, fmf);
680 Value absImagIsNotInf = arith::CmpFOp::create(
681 b, arith::CmpFPredicate::ONE, absImag, inf, fmf);
682 Value realIsInf =
683 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, real, inf, fmf);
684 Value realIsNegInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
685 real, negInf, fmf);
686
687 resultReal = arith::SelectOp::create(
688 b, arith::AndIOp::create(b, realIsNegInf, absImagIsNotInf), zero,
689 resultReal);
690 resultReal = arith::SelectOp::create(
691 b, arith::OrIOp::create(b, absImagIsInf, realIsInf), inf, resultReal);
692
693 Value imagSignInf = math::CopySignOp::create(b, inf, imag, fmf);
694 resultImag = arith::SelectOp::create(
695 b,
696 arith::CmpFOp::create(b, arith::CmpFPredicate::UNO, absSqrt, absSqrt),
697 nan, resultImag);
698 resultImag = arith::SelectOp::create(
699 b, arith::OrIOp::create(b, absImagIsInf, realIsNegInf), imagSignInf,
700 resultImag);
701 }
702
703 Value resultIsZero =
704 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, absSqrt, zero, fmf);
705 resultReal = arith::SelectOp::create(b, resultIsZero, zero, resultReal);
706 resultImag = arith::SelectOp::create(b, resultIsZero, zero, resultImag);
707
708 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
709 resultImag);
710 return success();
711 }
712};
713
714struct SignOpConversion : public OpConversionPattern<complex::SignOp> {
715 using OpConversionPattern<complex::SignOp>::OpConversionPattern;
716
717 LogicalResult
718 matchAndRewrite(complex::SignOp op, OpAdaptor adaptor,
719 ConversionPatternRewriter &rewriter) const override {
720 auto type = cast<ComplexType>(adaptor.getComplex().getType());
721 auto elementType = cast<FloatType>(type.getElementType());
722 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
723 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
724
725 Value real = complex::ReOp::create(b, elementType, adaptor.getComplex());
726 Value imag = complex::ImOp::create(b, elementType, adaptor.getComplex());
727 Value zero =
728 arith::ConstantOp::create(b, elementType, b.getZeroAttr(elementType));
729 Value realIsZero =
730 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, real, zero);
731 Value imagIsZero =
732 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, imag, zero);
733 Value isZero = arith::AndIOp::create(b, realIsZero, imagIsZero);
734 auto abs =
735 complex::AbsOp::create(b, elementType, adaptor.getComplex(), fmf);
736 Value realSign = arith::DivFOp::create(b, real, abs, fmf);
737 Value imagSign = arith::DivFOp::create(b, imag, abs, fmf);
738 Value sign = complex::CreateOp::create(b, type, realSign, imagSign);
739 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, isZero,
740 adaptor.getComplex(), sign);
741 return success();
742 }
743};
744
745template <typename Op>
746struct TanTanhOpConversion : public OpConversionPattern<Op> {
747 using OpConversionPattern<Op>::OpConversionPattern;
748
749 LogicalResult
750 matchAndRewrite(Op op, typename Op::Adaptor adaptor,
751 ConversionPatternRewriter &rewriter) const override {
752 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
753 auto loc = op.getLoc();
754 auto type = cast<ComplexType>(adaptor.getComplex().getType());
755 auto elementType = cast<FloatType>(type.getElementType());
756 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
757 const auto &floatSemantics = elementType.getFloatSemantics();
758
759 Value real =
760 complex::ReOp::create(b, loc, elementType, adaptor.getComplex());
761 Value imag =
762 complex::ImOp::create(b, loc, elementType, adaptor.getComplex());
763
764 if constexpr (std::is_same_v<Op, complex::TanOp>) {
765 // tan(x+yi) = -i*tanh(-y + xi)
766 std::swap(real, imag);
767 real = arith::NegFOp::create(b, real, fmf);
768 }
769
770 auto cst = [&](APFloat v) {
771 return arith::ConstantOp::create(b, elementType,
772 b.getFloatAttr(elementType, v));
773 };
774 Value inf = cst(APFloat::getInf(floatSemantics));
775 Value four = arith::ConstantOp::create(b, elementType,
776 b.getFloatAttr(elementType, 4.0));
777 Value twoReal = arith::AddFOp::create(b, real, real, fmf);
778 Value negTwoReal = arith::NegFOp::create(b, twoReal, fmf);
779 Value expTwoRealMinusOne = math::ExpM1Op::create(b, twoReal, fmf);
780 Value expNegTwoRealMinusOne = math::ExpM1Op::create(b, negTwoReal, fmf);
781 Value realNum = arith::SubFOp::create(b, expTwoRealMinusOne,
782 expNegTwoRealMinusOne, fmf);
783 Value expProduct = arith::MulFOp::create(b, expTwoRealMinusOne,
784 expNegTwoRealMinusOne, fmf);
785 Value expSumMinusTwo = arith::NegFOp::create(b, expProduct, fmf);
786
787 Value cosImag = math::CosOp::create(b, imag, fmf);
788 Value cosImagSq = arith::MulFOp::create(b, cosImag, cosImag, fmf);
789 Value twoCosTwoImagPlusOne = arith::MulFOp::create(b, cosImagSq, four, fmf);
790 Value sinImag = math::SinOp::create(b, imag, fmf);
791
792 Value imagNum = arith::MulFOp::create(
793 b, four, arith::MulFOp::create(b, cosImag, sinImag, fmf), fmf);
794
795 Value denom =
796 arith::AddFOp::create(b, expSumMinusTwo, twoCosTwoImagPlusOne, fmf);
797
798 Value isInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
799 expSumMinusTwo, inf, fmf);
800 Value negOne = arith::ConstantOp::create(b, elementType,
801 b.getFloatAttr(elementType, -1.0));
802 Value realLimit = math::CopySignOp::create(b, negOne, real, fmf);
803
804 Value resultReal = arith::SelectOp::create(
805 b, isInf, realLimit, arith::DivFOp::create(b, realNum, denom, fmf));
806 Value resultImag = arith::DivFOp::create(b, imagNum, denom, fmf);
807
808 if (!arith::bitEnumContainsAll(fmf, arith::FastMathFlags::nnan |
809 arith::FastMathFlags::ninf)) {
810 Value absReal = math::AbsFOp::create(b, real, fmf);
811 Value zero = arith::ConstantOp::create(b, elementType,
812 b.getFloatAttr(elementType, 0.0));
813 Value nan = cst(APFloat::getNaN(floatSemantics));
814
815 Value absRealIsInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
816 absReal, inf, fmf);
817 Value imagIsZero =
818 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, imag, zero, fmf);
819 Value absRealIsNotInf = arith::XOrIOp::create(
820 b, absRealIsInf, arith::ConstantIntOp::create(b, true, /*width=*/1));
821
822 Value imagNumIsNaN = arith::CmpFOp::create(b, arith::CmpFPredicate::UNO,
823 imagNum, imagNum, fmf);
824 Value resultRealIsNaN =
825 arith::AndIOp::create(b, imagNumIsNaN, absRealIsNotInf);
826 Value resultImagIsZero = arith::OrIOp::create(
827 b, imagIsZero, arith::AndIOp::create(b, absRealIsInf, imagNumIsNaN));
828
829 resultReal = arith::SelectOp::create(b, resultRealIsNaN, nan, resultReal);
830 resultImag =
831 arith::SelectOp::create(b, resultImagIsZero, zero, resultImag);
832 }
833
834 if constexpr (std::is_same_v<Op, complex::TanOp>) {
835 // tan(x+yi) = -i*tanh(-y + xi)
836 std::swap(resultReal, resultImag);
837 resultImag = arith::NegFOp::create(b, resultImag, fmf);
838 }
839
840 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
841 resultImag);
842 return success();
843 }
844};
845
846struct ConjOpConversion : public OpConversionPattern<complex::ConjOp> {
847 using OpConversionPattern<complex::ConjOp>::OpConversionPattern;
848
849 LogicalResult
850 matchAndRewrite(complex::ConjOp op, OpAdaptor adaptor,
851 ConversionPatternRewriter &rewriter) const override {
852 auto loc = op.getLoc();
853 auto type = cast<ComplexType>(adaptor.getComplex().getType());
854 auto elementType = cast<FloatType>(type.getElementType());
855 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
856 Value real =
857 complex::ReOp::create(rewriter, loc, elementType, adaptor.getComplex());
858 Value imag =
859 complex::ImOp::create(rewriter, loc, elementType, adaptor.getComplex());
860 Value negImag =
861 arith::NegFOp::create(rewriter, loc, elementType, imag, fmf);
862
863 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, real, negImag);
864
865 return success();
866 }
867};
868
869/// Converts lhs^y = (a+bi)^(c+di) to
870/// (a*a+b*b)^(0.5c) * exp(-d*atan2(b,a)) * (cos(q) + i*sin(q)),
871/// where q = c*atan2(b,a)+0.5d*ln(a*a+b*b)
872static Value powOpConversionImpl(mlir::ImplicitLocOpBuilder &builder,
873 ComplexType type, Value lhs, Value c, Value d,
874 arith::FastMathFlags fmf) {
875 auto elementType = cast<FloatType>(type.getElementType());
876
877 Value a = complex::ReOp::create(builder, lhs);
878 Value b = complex::ImOp::create(builder, lhs);
879
880 Value abs = complex::AbsOp::create(builder, lhs, fmf);
881 Value absToC = math::PowFOp::create(builder, abs, c, fmf);
882
883 Value negD = arith::NegFOp::create(builder, d, fmf);
884 Value argLhs = math::Atan2Op::create(builder, b, a, fmf);
885 Value negDArgLhs = arith::MulFOp::create(builder, negD, argLhs, fmf);
886 Value expNegDArgLhs = math::ExpOp::create(builder, negDArgLhs, fmf);
887
888 Value coeff = arith::MulFOp::create(builder, absToC, expNegDArgLhs, fmf);
889 Value lnAbs = math::LogOp::create(builder, abs, fmf);
890 Value cArgLhs = arith::MulFOp::create(builder, c, argLhs, fmf);
891 Value dLnAbs = arith::MulFOp::create(builder, d, lnAbs, fmf);
892 Value q = arith::AddFOp::create(builder, cArgLhs, dLnAbs, fmf);
893 Value cosQ = math::CosOp::create(builder, q, fmf);
894 Value sinQ = math::SinOp::create(builder, q, fmf);
895
896 Value inf = arith::ConstantOp::create(
897 builder, elementType,
898 builder.getFloatAttr(elementType,
899 APFloat::getInf(elementType.getFloatSemantics())));
900 Value zero = arith::ConstantOp::create(
901 builder, elementType, builder.getFloatAttr(elementType, 0.0));
902 Value one = arith::ConstantOp::create(builder, elementType,
903 builder.getFloatAttr(elementType, 1.0));
904 Value complexOne = complex::CreateOp::create(builder, type, one, zero);
905 Value complexZero = complex::CreateOp::create(builder, type, zero, zero);
906 Value complexInf = complex::CreateOp::create(builder, type, inf, zero);
907
908 // Case 0:
909 // d^c is 0 if d is 0 and c > 0. 0^0 is defined to be 1.0, see
910 // Branch Cuts for Complex Elementary Functions or Much Ado About
911 // Nothing's Sign Bit, W. Kahan, Section 10.
912 Value absEqZero =
913 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, abs, zero, fmf);
914 Value dEqZero =
915 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, d, zero, fmf);
916 Value cEqZero =
917 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, c, zero, fmf);
918 Value bEqZero =
919 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, b, zero, fmf);
920
921 Value zeroLeC =
922 arith::CmpFOp::create(builder, arith::CmpFPredicate::OLE, zero, c, fmf);
923 Value coeffCosQ = arith::MulFOp::create(builder, coeff, cosQ, fmf);
924 Value coeffSinQ = arith::MulFOp::create(builder, coeff, sinQ, fmf);
925 Value complexOneOrZero =
926 arith::SelectOp::create(builder, cEqZero, complexOne, complexZero);
927 Value coeffCosSin =
928 complex::CreateOp::create(builder, type, coeffCosQ, coeffSinQ);
929 Value cutoff0 = arith::SelectOp::create(
930 builder,
931 arith::AndIOp::create(
932 builder, arith::AndIOp::create(builder, absEqZero, dEqZero), zeroLeC),
933 complexOneOrZero, coeffCosSin);
934
935 // Case 1:
936 // x^0 is defined to be 1 for any x, see
937 // Branch Cuts for Complex Elementary Functions or Much Ado About
938 // Nothing's Sign Bit, W. Kahan, Section 10.
939 Value rhsEqZero = arith::AndIOp::create(builder, cEqZero, dEqZero);
940 Value cutoff1 =
941 arith::SelectOp::create(builder, rhsEqZero, complexOne, cutoff0);
942
943 // Case 2:
944 // 1^(c + d*i) = 1 + 0*i
945 Value lhsEqOne = arith::AndIOp::create(
946 builder,
947 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, a, one, fmf),
948 bEqZero);
949 Value cutoff2 =
950 arith::SelectOp::create(builder, lhsEqOne, complexOne, cutoff1);
951
952 // Case 3:
953 // inf^(c + 0*i) = inf + 0*i, c > 0
954 Value lhsEqInf = arith::AndIOp::create(
955 builder,
956 arith::CmpFOp::create(builder, arith::CmpFPredicate::OEQ, a, inf, fmf),
957 bEqZero);
958 Value rhsGt0 = arith::AndIOp::create(
959 builder, dEqZero,
960 arith::CmpFOp::create(builder, arith::CmpFPredicate::OGT, c, zero, fmf));
961 Value cutoff3 = arith::SelectOp::create(
962 builder, arith::AndIOp::create(builder, lhsEqInf, rhsGt0), complexInf,
963 cutoff2);
964
965 // Case 4:
966 // inf^(c + 0*i) = 0 + 0*i, c < 0
967 Value rhsLt0 = arith::AndIOp::create(
968 builder, dEqZero,
969 arith::CmpFOp::create(builder, arith::CmpFPredicate::OLT, c, zero, fmf));
970 Value cutoff4 = arith::SelectOp::create(
971 builder, arith::AndIOp::create(builder, lhsEqInf, rhsLt0), complexZero,
972 cutoff3);
973
974 return cutoff4;
975}
976
977struct PowiOpConversion : public OpConversionPattern<complex::PowiOp> {
978 using OpConversionPattern<complex::PowiOp>::OpConversionPattern;
979
980 LogicalResult
981 matchAndRewrite(complex::PowiOp op, OpAdaptor adaptor,
982 ConversionPatternRewriter &rewriter) const override {
983 ImplicitLocOpBuilder builder(op.getLoc(), rewriter);
984 auto type = cast<ComplexType>(op.getType());
985 auto elementType = cast<FloatType>(type.getElementType());
986
987 Value floatExponent =
988 arith::SIToFPOp::create(builder, elementType, adaptor.getRhs());
989 Value zero = arith::ConstantOp::create(
990 builder, elementType, builder.getFloatAttr(elementType, 0.0));
991 Value complexExponent =
992 complex::CreateOp::create(builder, type, floatExponent, zero);
993
994 auto pow = complex::PowOp::create(builder, type, adaptor.getLhs(),
995 complexExponent, op.getFastmathAttr());
996 rewriter.replaceOp(op, pow.getResult());
997 return success();
998 }
999};
1000
1001struct PowOpConversion : public OpConversionPattern<complex::PowOp> {
1002 using OpConversionPattern<complex::PowOp>::OpConversionPattern;
1003
1004 LogicalResult
1005 matchAndRewrite(complex::PowOp op, OpAdaptor adaptor,
1006 ConversionPatternRewriter &rewriter) const override {
1007 mlir::ImplicitLocOpBuilder builder(op.getLoc(), rewriter);
1008 auto type = cast<ComplexType>(adaptor.getLhs().getType());
1009 auto elementType = cast<FloatType>(type.getElementType());
1010
1011 Value c = complex::ReOp::create(builder, elementType, adaptor.getRhs());
1012 Value d = complex::ImOp::create(builder, elementType, adaptor.getRhs());
1013
1014 rewriter.replaceOp(op, {powOpConversionImpl(builder, type, adaptor.getLhs(),
1015 c, d, op.getFastmath())});
1016 return success();
1017 }
1018};
1019
1020struct RsqrtOpConversion : public OpConversionPattern<complex::RsqrtOp> {
1021 using OpConversionPattern<complex::RsqrtOp>::OpConversionPattern;
1022
1023 LogicalResult
1024 matchAndRewrite(complex::RsqrtOp op, OpAdaptor adaptor,
1025 ConversionPatternRewriter &rewriter) const override {
1026 mlir::ImplicitLocOpBuilder b(op.getLoc(), rewriter);
1027 auto type = cast<ComplexType>(adaptor.getComplex().getType());
1028 auto elementType = cast<FloatType>(type.getElementType());
1029
1030 arith::FastMathFlags fmf = op.getFastMathFlagsAttr().getValue();
1031
1032 auto cst = [&](APFloat v) {
1033 return arith::ConstantOp::create(b, elementType,
1034 b.getFloatAttr(elementType, v));
1035 };
1036 const auto &floatSemantics = elementType.getFloatSemantics();
1037 Value zero = cst(APFloat::getZero(floatSemantics));
1038 Value inf = cst(APFloat::getInf(floatSemantics));
1039 Value negHalf = arith::ConstantOp::create(
1040 b, elementType, b.getFloatAttr(elementType, -0.5));
1041 Value nan = cst(APFloat::getNaN(floatSemantics));
1042
1043 Value real = complex::ReOp::create(b, elementType, adaptor.getComplex());
1044 Value imag = complex::ImOp::create(b, elementType, adaptor.getComplex());
1045 Value absRsqrt = computeAbs(real, imag, fmf, b, AbsFn::rsqrt);
1046 Value argArg = math::Atan2Op::create(b, imag, real, fmf);
1047 Value rsqrtArg = arith::MulFOp::create(b, argArg, negHalf, fmf);
1048 Value cos = math::CosOp::create(b, rsqrtArg, fmf);
1049 Value sin = math::SinOp::create(b, rsqrtArg, fmf);
1050
1051 Value resultReal = arith::MulFOp::create(b, absRsqrt, cos, fmf);
1052 Value resultImag = arith::MulFOp::create(b, absRsqrt, sin, fmf);
1053
1054 if (!arith::bitEnumContainsAll(fmf, arith::FastMathFlags::nnan |
1055 arith::FastMathFlags::ninf)) {
1056 Value realSignedZero = math::CopySignOp::create(b, zero, real, fmf);
1057 Value imagSignedZero = math::CopySignOp::create(b, zero, imag, fmf);
1058 Value negImagSignedZero = arith::NegFOp::create(b, imagSignedZero, fmf);
1059
1060 Value absReal = math::AbsFOp::create(b, real, fmf);
1061 Value absImag = math::AbsFOp::create(b, imag, fmf);
1062
1063 Value absImagIsInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
1064 absImag, inf, fmf);
1065 Value realIsNan =
1066 arith::CmpFOp::create(b, arith::CmpFPredicate::UNO, real, real, fmf);
1067 Value realIsInf = arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ,
1068 absReal, inf, fmf);
1069 Value inIsNanInf = arith::AndIOp::create(b, absImagIsInf, realIsNan);
1070
1071 Value resultIsZero = arith::OrIOp::create(b, inIsNanInf, realIsInf);
1072
1073 resultReal =
1074 arith::SelectOp::create(b, resultIsZero, realSignedZero, resultReal);
1075 resultImag = arith::SelectOp::create(b, resultIsZero, negImagSignedZero,
1076 resultImag);
1077 }
1078
1079 Value isRealZero =
1080 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, real, zero, fmf);
1081 Value isImagZero =
1082 arith::CmpFOp::create(b, arith::CmpFPredicate::OEQ, imag, zero, fmf);
1083 Value isZero = arith::AndIOp::create(b, isRealZero, isImagZero);
1084
1085 resultReal = arith::SelectOp::create(b, isZero, inf, resultReal);
1086 resultImag = arith::SelectOp::create(b, isZero, nan, resultImag);
1087
1088 rewriter.replaceOpWithNewOp<complex::CreateOp>(op, type, resultReal,
1089 resultImag);
1090 return success();
1091 }
1092};
1093
1094struct AngleOpConversion : public OpConversionPattern<complex::AngleOp> {
1095 using OpConversionPattern<complex::AngleOp>::OpConversionPattern;
1096
1097 LogicalResult
1098 matchAndRewrite(complex::AngleOp op, OpAdaptor adaptor,
1099 ConversionPatternRewriter &rewriter) const override {
1100 auto loc = op.getLoc();
1101 auto type = op.getType();
1102 arith::FastMathFlagsAttr fmf = op.getFastMathFlagsAttr();
1103
1104 Value real =
1105 complex::ReOp::create(rewriter, loc, type, adaptor.getComplex());
1106 Value imag =
1107 complex::ImOp::create(rewriter, loc, type, adaptor.getComplex());
1108
1109 rewriter.replaceOpWithNewOp<math::Atan2Op>(op, imag, real, fmf);
1110
1111 return success();
1112 }
1113};
1114
1115} // namespace
1116
1118 RewritePatternSet &patterns, complex::ComplexRangeFlags complexRange) {
1119 // clang-format off
1120 patterns.add<
1121 AbsOpConversion,
1122 AngleOpConversion,
1123 Atan2OpConversion,
1124 BinaryComplexOpConversion<complex::AddOp, arith::AddFOp>,
1125 BinaryComplexOpConversion<complex::SubOp, arith::SubFOp>,
1126 ComparisonOpConversion<complex::EqualOp, arith::CmpFPredicate::OEQ>,
1127 ComparisonOpConversion<complex::NotEqualOp, arith::CmpFPredicate::UNE>,
1128 ConjOpConversion,
1129 CosOpConversion,
1130 ExpOpConversion,
1131 Expm1OpConversion,
1132 Log1pOpConversion,
1133 LogOpConversion,
1134 MulOpConversion,
1135 NegOpConversion,
1136 SignOpConversion,
1137 SinOpConversion,
1138 SqrtOpConversion,
1139 TanTanhOpConversion<complex::TanOp>,
1140 TanTanhOpConversion<complex::TanhOp>,
1141 PowiOpConversion,
1142 PowOpConversion,
1143 RsqrtOpConversion
1144 >(patterns.getContext());
1145
1146 patterns.add<DivOpConversion>(patterns.getContext(), complexRange);
1147
1148 // clang-format on
1149}
1150
1151namespace {
1152struct ConvertComplexToStandardPass
1153 : public impl::ConvertComplexToStandardPassBase<
1154 ConvertComplexToStandardPass> {
1155 using Base::Base;
1156
1157 void runOnOperation() override;
1158};
1159
1160void ConvertComplexToStandardPass::runOnOperation() {
1161 // Convert to the Standard dialect using the converter defined above.
1162 RewritePatternSet patterns(&getContext());
1163 populateComplexToStandardConversionPatterns(patterns, complexRange);
1164
1165 ConversionTarget target(getContext());
1166 target.addLegalDialect<arith::ArithDialect, math::MathDialect>();
1167 target.addLegalOp<complex::CreateOp, complex::ImOp, complex::ReOp>();
1168 if (failed(
1169 applyPartialConversion(getOperation(), target, std::move(patterns))))
1170 signalPassFailure();
1171}
1172} // namespace
return success()
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Definition Builders.h:632
Location getLoc()
The source location the operation was defined or derived from.
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.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
Definition ArithOps.cpp:297
NestedPattern Op(FilterFunctionType filter=defaultFilterFunction)
void convertDivToStandardUsingAlgebraic(ConversionPatternRewriter &rewriter, Location loc, Value lhsRe, Value lhsIm, Value rhsRe, Value rhsIm, arith::FastMathFlagsAttr fmf, Value *resultRe, Value *resultIm)
convert a complex division to the arith/math dialects using algebraic method
void convertDivToStandardUsingRangeReduction(ConversionPatternRewriter &rewriter, Location loc, Value lhsRe, Value lhsIm, Value rhsRe, Value rhsIm, arith::FastMathFlagsAttr fmf, Value *resultRe, Value *resultIm)
convert a complex division to the arith/math dialects using Smith's method
Fraction abs(const Fraction &f)
Definition Fraction.h:107
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:733
OwningOpRef< spirv::ModuleOp > combine(ArrayRef< spirv::ModuleOp > inputModules, OpBuilder &combinedModuleBuilder, SymbolRenameListener symRenameListener)
Combines a list of SPIR-V inputModules into one.
Include the generated interface declarations.
constexpr T real(const NonFloatComplex< T > &x)
Definition Complex.h:255
void populateComplexToStandardConversionPatterns(RewritePatternSet &patterns, mlir::complex::ComplexRangeFlags complexRange=mlir::complex::ComplexRangeFlags::improved)
Populate the given list with patterns that convert from Complex to Standard.
constexpr T imag(const NonFloatComplex< T > &x)
Definition Complex.h:260