MLIR 24.0.0git
MathToLLVM.cpp
Go to the documentation of this file.
1//===- MathToLLVM.cpp - Math to LLVM dialect conversion -------------------===//
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
18#include "mlir/IR/Matchers.h"
20#include "mlir/Pass/Pass.h"
21
22#include "llvm/ADT/FloatingPointMode.h"
23
24namespace mlir {
25#define GEN_PASS_DEF_CONVERTMATHTOLLVMPASS
26#include "mlir/Conversion/Passes.h.inc"
27} // namespace mlir
28
29using namespace mlir;
30
31namespace {
32
33template <typename SourceOp, typename TargetOp>
35
36template <typename SourceOp, typename TargetOp, bool FailOnUnsupportedFP = true>
37using ConvertFMFMathToLLVMPattern =
38 VectorConvertToLLVMPattern<SourceOp, TargetOp, ConvertFastMath,
39 FailOnUnsupportedFP>;
40
41/// Lowering pattern that matches only when the source op's rounding mode
42/// presence agrees with `HasRoundingMode`. Mirrors the helper of the same
43/// name in `mlir/lib/Conversion/ArithToLLVM/ArithToLLVM.cpp`. This lets us
44/// register two patterns for one math op: an unconstrained one that lowers
45/// to a regular LLVM op, and a constrained one (rounding mode present) that
46/// lowers to an `llvm.intr.experimental.constrained.*` intrinsic.
47template <typename SourceOp, typename TargetOp, bool HasRoundingMode,
48 template <typename, typename> typename AttrConvert =
50 bool FailOnUnsupportedFP = true>
51struct ConstrainedVectorConvertToLLVMPattern
52 : public VectorConvertToLLVMPattern<SourceOp, TargetOp, AttrConvert,
53 FailOnUnsupportedFP> {
54 using VectorConvertToLLVMPattern<
55 SourceOp, TargetOp, AttrConvert,
56 FailOnUnsupportedFP>::VectorConvertToLLVMPattern;
57
58 LogicalResult
59 matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
60 ConversionPatternRewriter &rewriter) const override {
61 if (HasRoundingMode != static_cast<bool>(op.getRoundingModeAttr()))
62 return failure();
63 return VectorConvertToLLVMPattern<
64 SourceOp, TargetOp, AttrConvert,
65 FailOnUnsupportedFP>::matchAndRewrite(op, adaptor, rewriter);
66 }
67};
68
69using AbsFOpLowering =
70 ConvertFMFMathToLLVMPattern<math::AbsFOp, LLVM::FAbsOp,
71 /*FailOnUnsupportedFP=*/true>;
72using CeilOpLowering = ConvertFMFMathToLLVMPattern<math::CeilOp, LLVM::FCeilOp>;
73using CopySignOpLowering =
74 ConvertFMFMathToLLVMPattern<math::CopySignOp, LLVM::CopySignOp>;
75using CosOpLowering = ConvertFMFMathToLLVMPattern<math::CosOp, LLVM::CosOp>;
76using CoshOpLowering = ConvertFMFMathToLLVMPattern<math::CoshOp, LLVM::CoshOp>;
77using AcosOpLowering = ConvertFMFMathToLLVMPattern<math::AcosOp, LLVM::ACosOp>;
78using CtPopFOpLowering =
79 VectorConvertToLLVMPattern<math::CtPopOp, LLVM::CtPopOp,
81 /*FailOnUnsupportedFP=*/true>;
82using Exp2OpLowering = ConvertFMFMathToLLVMPattern<math::Exp2Op, LLVM::Exp2Op>;
83using ExpOpLowering = ConvertFMFMathToLLVMPattern<math::ExpOp, LLVM::ExpOp>;
84using FloorOpLowering =
85 ConvertFMFMathToLLVMPattern<math::FloorOp, LLVM::FFloorOp>;
86using FmaOpLowering =
87 ConstrainedVectorConvertToLLVMPattern<math::FmaOp, LLVM::FMAOp,
88 /*HasRoundingMode=*/false,
89 ConvertFastMath,
90 /*FailOnUnsupportedFP=*/true>;
91using ConstrainedFmaOpLowering = ConstrainedVectorConvertToLLVMPattern<
92 math::FmaOp, LLVM::ConstrainedFMAIntr, /*HasRoundingMode=*/true,
93 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
94using Log10OpLowering =
95 ConvertFMFMathToLLVMPattern<math::Log10Op, LLVM::Log10Op>;
96using Log2OpLowering = ConvertFMFMathToLLVMPattern<math::Log2Op, LLVM::Log2Op>;
97using LogOpLowering = ConvertFMFMathToLLVMPattern<math::LogOp, LLVM::LogOp>;
98using PowFOpLowering = ConvertFMFMathToLLVMPattern<math::PowFOp, LLVM::PowOp>;
99using RoundEvenOpLowering =
100 ConvertFMFMathToLLVMPattern<math::RoundEvenOp, LLVM::RoundEvenOp>;
101using RoundOpLowering =
102 ConvertFMFMathToLLVMPattern<math::RoundOp, LLVM::RoundOp>;
103using SinOpLowering = ConvertFMFMathToLLVMPattern<math::SinOp, LLVM::SinOp>;
104using SinhOpLowering = ConvertFMFMathToLLVMPattern<math::SinhOp, LLVM::SinhOp>;
105using ASinOpLowering = ConvertFMFMathToLLVMPattern<math::AsinOp, LLVM::ASinOp>;
106using SqrtOpLowering = ConvertFMFMathToLLVMPattern<math::SqrtOp, LLVM::SqrtOp>;
107using FTruncOpLowering =
108 ConvertFMFMathToLLVMPattern<math::TruncOp, LLVM::FTruncOp>;
109using TanOpLowering = ConvertFMFMathToLLVMPattern<math::TanOp, LLVM::TanOp>;
110using TanhOpLowering = ConvertFMFMathToLLVMPattern<math::TanhOp, LLVM::TanhOp>;
111using ATanOpLowering = ConvertFMFMathToLLVMPattern<math::AtanOp, LLVM::ATanOp>;
112using ATan2OpLowering =
113 ConvertFMFMathToLLVMPattern<math::Atan2Op, LLVM::ATan2Op>;
114// A `CtLz/CtTz/absi(a)` is converted into `CtLz/CtTz/absi(a, false)`.
115// TODO: Result and operand types match for `absi` as opposed to `ct*z`, so it
116// may be better to separate the patterns.
117template <typename MathOp, typename LLVMOp>
118struct IntOpWithFlagLowering
119 : public ConvertOpToLLVMPattern<MathOp, /*FailOnUnsupportedFP=*/true> {
120 using ConvertOpToLLVMPattern<
121 MathOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
122 using Super = IntOpWithFlagLowering<MathOp, LLVMOp>;
123
124 LogicalResult
125 matchAndRewrite(MathOp op, typename MathOp::Adaptor adaptor,
126 ConversionPatternRewriter &rewriter) const override {
127 const auto &typeConverter = *this->getTypeConverter();
128 auto operandType = adaptor.getOperand().getType();
129 auto llvmOperandType = typeConverter.convertType(operandType);
130 if (!llvmOperandType)
131 return failure();
132
133 auto loc = op.getLoc();
134 auto resultType = op.getResult().getType();
135 auto llvmResultType = typeConverter.convertType(resultType);
136 if (!llvmResultType)
137 return failure();
138
139 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
140 rewriter.replaceOpWithNewOp<LLVMOp>(op, llvmResultType,
141 adaptor.getOperand(), false);
142 return success();
143 }
144
145 if (!isa<VectorType>(resultType))
146 return failure();
147
149 op.getOperation(), adaptor.getOperands(), typeConverter,
150 [&](Type llvm1DVectorTy, ValueRange operands) {
151 return LLVMOp::create(rewriter, loc, llvm1DVectorTy, operands[0],
152 false);
153 },
154 rewriter);
155 }
156};
157
158using CountLeadingZerosOpLowering =
159 IntOpWithFlagLowering<math::CountLeadingZerosOp, LLVM::CountLeadingZerosOp>;
160using CountTrailingZerosOpLowering =
161 IntOpWithFlagLowering<math::CountTrailingZerosOp,
162 LLVM::CountTrailingZerosOp>;
163using AbsIOpLowering = IntOpWithFlagLowering<math::AbsIOp, LLVM::AbsOp>;
164
165// A `sincos` is converted into `llvm.intr.sincos` followed by extractvalue ops.
166struct SincosOpLowering
167 : public ConvertOpToLLVMPattern<math::SincosOp,
168 /*FailOnUnsupportedFP=*/true> {
170 math::SincosOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
171
172 LogicalResult
173 matchAndRewrite(math::SincosOp op, OpAdaptor adaptor,
174 ConversionPatternRewriter &rewriter) const override {
175 const LLVMTypeConverter &typeConverter = *this->getTypeConverter();
176 mlir::Location loc = op.getLoc();
177 mlir::Type operandType = adaptor.getOperand().getType();
178 mlir::Type llvmOperandType = typeConverter.convertType(operandType);
179 mlir::Type sinType = typeConverter.convertType(op.getSin().getType());
180 mlir::Type cosType = typeConverter.convertType(op.getCos().getType());
181 if (!llvmOperandType || !sinType || !cosType)
182 return failure();
183
184 ConvertFastMath<math::SincosOp, LLVM::SincosOp> attrs(op);
185
186 auto structType = LLVM::LLVMStructType::getLiteral(
187 rewriter.getContext(), {llvmOperandType, llvmOperandType});
188
189 auto sincosOp = LLVM::SincosOp::create(
190 rewriter, loc, TypeRange{structType}, ValueRange{adaptor.getOperand()},
191 attrs.getProperties(), attrs.getDiscardableAttrs());
192
193 auto sinValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 0);
194 auto cosValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 1);
195
196 rewriter.replaceOp(op, {sinValue, cosValue});
197 return success();
198 }
199};
200
201// A `expm1` is converted into `exp - 1`.
202struct ExpM1OpLowering
203 : public ConvertOpToLLVMPattern<math::ExpM1Op,
204 /*FailOnUnsupportedFP=*/true> {
205 using ConvertOpToLLVMPattern<
206 math::ExpM1Op, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
207
208 LogicalResult
209 matchAndRewrite(math::ExpM1Op op, OpAdaptor adaptor,
210 ConversionPatternRewriter &rewriter) const override {
211 const auto &typeConverter = *this->getTypeConverter();
212 auto operandType = adaptor.getOperand().getType();
213 auto llvmOperandType = typeConverter.convertType(operandType);
214 if (!llvmOperandType)
215 return failure();
216
217 auto loc = op.getLoc();
218 auto resultType = op.getResult().getType();
219 auto floatType = cast<FloatType>(
220 typeConverter.convertType(getElementTypeOrSelf(resultType)));
221 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
222 ConvertFastMath<math::ExpM1Op, LLVM::ExpOp> expAttrs(op);
223 ConvertFastMath<math::ExpM1Op, LLVM::FSubOp> subAttrs(op);
224
225 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
226 LLVM::ConstantOp one;
227 if (LLVM::isCompatibleVectorType(llvmOperandType)) {
228 one = LLVM::ConstantOp::create(
229 rewriter, loc, llvmOperandType,
230 SplatElementsAttr::get(cast<ShapedType>(llvmOperandType),
231 floatOne));
232 } else {
233 one =
234 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
235 }
236 auto exp = LLVM::ExpOp::create(rewriter, loc, TypeRange{llvmOperandType},
237 ValueRange{adaptor.getOperand()},
238 expAttrs.getProperties(),
239 expAttrs.getDiscardableAttrs());
240 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(
241 op, TypeRange{llvmOperandType}, ValueRange{exp, one},
242 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
243 return success();
244 }
245
246 if (!isa<VectorType>(resultType))
247 return rewriter.notifyMatchFailure(op, "expected vector result type");
248
250 op.getOperation(), adaptor.getOperands(), typeConverter,
251 [&](Type llvm1DVectorTy, ValueRange operands) {
252 auto numElements = LLVM::getVectorNumElements(llvm1DVectorTy);
253 auto splatAttr = SplatElementsAttr::get(
254 mlir::VectorType::get({numElements.getKnownMinValue()}, floatType,
255 {numElements.isScalable()}),
256 floatOne);
257 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
258 splatAttr);
259 auto exp = LLVM::ExpOp::create(
260 rewriter, loc, TypeRange{llvm1DVectorTy}, ValueRange{operands[0]},
261 expAttrs.getProperties(), expAttrs.getDiscardableAttrs());
262 return LLVM::FSubOp::create(
263 rewriter, loc, TypeRange{llvm1DVectorTy}, ValueRange{exp, one},
264 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
265 },
266 rewriter);
267 }
268};
269
270// A `log1p` is converted into `log(1 + ...)`.
271struct Log1pOpLowering
272 : public ConvertOpToLLVMPattern<math::Log1pOp,
273 /*FailOnUnsupportedFP=*/true> {
274 using ConvertOpToLLVMPattern<
275 math::Log1pOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
276
277 LogicalResult
278 matchAndRewrite(math::Log1pOp op, OpAdaptor adaptor,
279 ConversionPatternRewriter &rewriter) const override {
280 const auto &typeConverter = *this->getTypeConverter();
281 auto operandType = adaptor.getOperand().getType();
282 auto llvmOperandType = typeConverter.convertType(operandType);
283 if (!llvmOperandType)
284 return rewriter.notifyMatchFailure(op, "unsupported operand type");
285
286 auto loc = op.getLoc();
287 auto resultType = op.getResult().getType();
288 auto floatType = cast<FloatType>(
289 typeConverter.convertType(getElementTypeOrSelf(resultType)));
290 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
291 ConvertFastMath<math::Log1pOp, LLVM::FAddOp> addAttrs(op);
292 ConvertFastMath<math::Log1pOp, LLVM::LogOp> logAttrs(op);
293
294 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
295 LLVM::ConstantOp one =
296 isa<VectorType>(llvmOperandType)
297 ? LLVM::ConstantOp::create(
298 rewriter, loc, llvmOperandType,
299 SplatElementsAttr::get(cast<ShapedType>(llvmOperandType),
300 floatOne))
301 : LLVM::ConstantOp::create(rewriter, loc, llvmOperandType,
302 floatOne);
303
304 auto add = LLVM::FAddOp::create(rewriter, loc, TypeRange{llvmOperandType},
305 ValueRange{one, adaptor.getOperand()},
306 addAttrs.getProperties(),
307 addAttrs.getDiscardableAttrs());
308 rewriter.replaceOpWithNewOp<LLVM::LogOp>(
309 op, TypeRange{llvmOperandType}, ValueRange{add},
310 logAttrs.getProperties(), logAttrs.getDiscardableAttrs());
311 return success();
312 }
313
314 if (!isa<VectorType>(resultType))
315 return rewriter.notifyMatchFailure(op, "expected vector result type");
316
318 op.getOperation(), adaptor.getOperands(), typeConverter,
319 [&](Type llvm1DVectorTy, ValueRange operands) {
320 auto numElements = LLVM::getVectorNumElements(llvm1DVectorTy);
321 auto splatAttr = SplatElementsAttr::get(
322 mlir::VectorType::get({numElements.getKnownMinValue()}, floatType,
323 {numElements.isScalable()}),
324 floatOne);
325 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
326 splatAttr);
327 auto add = LLVM::FAddOp::create(
328 rewriter, loc, TypeRange{llvm1DVectorTy},
329 ValueRange{one, operands[0]}, addAttrs.getProperties(),
330 addAttrs.getDiscardableAttrs());
331 return LLVM::LogOp::create(rewriter, loc, TypeRange{llvm1DVectorTy},
332 ValueRange{add}, logAttrs.getProperties(),
333 logAttrs.getDiscardableAttrs());
334 },
335 rewriter);
336 }
337};
338
339// A `rsqrt` is converted into `1 / sqrt`.
340struct RsqrtOpLowering
341 : public ConvertOpToLLVMPattern<math::RsqrtOp,
342 /*FailOnUnsupportedFP=*/true> {
343 using ConvertOpToLLVMPattern<
344 math::RsqrtOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
345
346 LogicalResult
347 matchAndRewrite(math::RsqrtOp op, OpAdaptor adaptor,
348 ConversionPatternRewriter &rewriter) const override {
349 const auto &typeConverter = *this->getTypeConverter();
350 auto operandType = adaptor.getOperand().getType();
351 auto llvmOperandType = typeConverter.convertType(operandType);
352 if (!llvmOperandType)
353 return failure();
354
355 auto loc = op.getLoc();
356 auto resultType = op.getResult().getType();
357 auto floatType = cast<FloatType>(
358 typeConverter.convertType(getElementTypeOrSelf(resultType)));
359 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
360 ConvertFastMath<math::RsqrtOp, LLVM::SqrtOp> sqrtAttrs(op);
361 ConvertFastMath<math::RsqrtOp, LLVM::FDivOp> divAttrs(op);
362
363 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
364 LLVM::ConstantOp one;
365 if (isa<VectorType>(llvmOperandType)) {
366 one = LLVM::ConstantOp::create(
367 rewriter, loc, llvmOperandType,
368 SplatElementsAttr::get(cast<ShapedType>(llvmOperandType),
369 floatOne));
370 } else {
371 one =
372 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
373 }
374 auto sqrt = LLVM::SqrtOp::create(
375 rewriter, loc, TypeRange{llvmOperandType},
376 ValueRange{adaptor.getOperand()}, sqrtAttrs.getProperties(),
377 sqrtAttrs.getDiscardableAttrs());
378 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(
379 op, TypeRange{llvmOperandType}, ValueRange{one, sqrt},
380 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
381 return success();
382 }
383
384 if (!isa<VectorType>(resultType))
385 return failure();
386
388 op.getOperation(), adaptor.getOperands(), typeConverter,
389 [&](Type llvm1DVectorTy, ValueRange operands) {
390 auto numElements = LLVM::getVectorNumElements(llvm1DVectorTy);
391 auto splatAttr = SplatElementsAttr::get(
392 mlir::VectorType::get({numElements.getKnownMinValue()}, floatType,
393 {numElements.isScalable()}),
394 floatOne);
395 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
396 splatAttr);
397 auto sqrt = LLVM::SqrtOp::create(
398 rewriter, loc, TypeRange{llvm1DVectorTy}, ValueRange{operands[0]},
399 sqrtAttrs.getProperties(), sqrtAttrs.getDiscardableAttrs());
400 return LLVM::FDivOp::create(
401 rewriter, loc, TypeRange{llvm1DVectorTy}, ValueRange{one, sqrt},
402 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
403 },
404 rewriter);
405 }
406};
407
408struct FPowIOpLowering
409 : public ConvertOpToLLVMPattern<math::FPowIOp,
410 /*FailOnUnsupportedFP=*/true> {
411 using ConvertOpToLLVMPattern<
412 math::FPowIOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
413
414 LogicalResult
415 matchAndRewrite(math::FPowIOp op, OpAdaptor adaptor,
416 ConversionPatternRewriter &rewriter) const override {
417 const auto &typeConverter = *this->getTypeConverter();
418 auto llvmOperandType = typeConverter.convertType(op.getLhs().getType());
419 if (!llvmOperandType)
420 return failure();
421
422 auto loc = op.getLoc();
423 Value exponent = adaptor.getRhs();
424 if (isa<VectorType>(op.getRhs().getType())) {
425 SplatElementsAttr splatAttr;
426 if (!matchPattern(op.getRhs(), m_Constant(&splatAttr)))
427 return rewriter.notifyMatchFailure(op, "expected a splat exponent");
428
429 auto exponentType = typeConverter.convertType(
430 getElementTypeOrSelf(op.getRhs().getType()));
431 if (!exponentType)
432 return failure();
433 exponent = LLVM::ConstantOp::create(
434 rewriter, loc, exponentType,
435 rewriter.getIntegerAttr(exponentType,
436 splatAttr.getSplatValue<APInt>()));
437 }
438
439 ConvertFastMath<math::FPowIOp, LLVM::PowIOp> attrs(op);
440
441 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
442 rewriter.replaceOpWithNewOp<LLVM::PowIOp>(
443 op, TypeRange{llvmOperandType},
444 ValueRange{adaptor.getLhs(), exponent}, attrs.getProperties(),
445 attrs.getDiscardableAttrs());
446 return success();
447 }
448
450 op.getOperation(), ValueRange{adaptor.getLhs()}, typeConverter,
451 [&](Type llvm1DVectorTy, ValueRange operands) {
452 return LLVM::PowIOp::create(rewriter, loc, TypeRange{llvm1DVectorTy},
453 ValueRange{operands[0], exponent},
454 attrs.getProperties(),
455 attrs.getDiscardableAttrs());
456 },
457 rewriter);
458 }
459};
460
461struct IsNaNOpLowering
462 : public ConvertOpToLLVMPattern<math::IsNaNOp,
463 /*FailOnUnsupportedFP=*/true> {
464 using ConvertOpToLLVMPattern<
465 math::IsNaNOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
466
467 LogicalResult
468 matchAndRewrite(math::IsNaNOp op, OpAdaptor adaptor,
469 ConversionPatternRewriter &rewriter) const override {
470 const auto &typeConverter = *this->getTypeConverter();
471 auto operandType =
472 typeConverter.convertType(adaptor.getOperand().getType());
473 auto resultType = typeConverter.convertType(op.getResult().getType());
474 if (!operandType || !resultType)
475 return failure();
476
477 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
478 op, resultType, adaptor.getOperand(), llvm::fcNan);
479 return success();
480 }
481};
482
483struct IsFiniteOpLowering
484 : public ConvertOpToLLVMPattern<math::IsFiniteOp,
485 /*FailOnUnsupportedFP=*/true> {
486 using ConvertOpToLLVMPattern<
487 math::IsFiniteOp, /*FailOnUnsupportedFP=*/true>::ConvertOpToLLVMPattern;
488
489 LogicalResult
490 matchAndRewrite(math::IsFiniteOp op, OpAdaptor adaptor,
491 ConversionPatternRewriter &rewriter) const override {
492 const auto &typeConverter = *this->getTypeConverter();
493 auto operandType =
494 typeConverter.convertType(adaptor.getOperand().getType());
495 auto resultType = typeConverter.convertType(op.getResult().getType());
496 if (!operandType || !resultType)
497 return failure();
498
499 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
500 op, resultType, adaptor.getOperand(), llvm::fcFinite);
501 return success();
502 }
503};
504
505struct ConvertMathToLLVMPass
506 : public impl::ConvertMathToLLVMPassBase<ConvertMathToLLVMPass> {
507 using Base::Base;
508
509 void runOnOperation() override {
510 RewritePatternSet patterns(&getContext());
511 LLVMTypeConverter converter(&getContext());
512 populateMathToLLVMConversionPatterns(converter, patterns, approximateLog1p);
513 LLVMConversionTarget target(getContext());
514 if (failed(applyPartialConversion(getOperation(), target,
515 std::move(patterns))))
516 signalPassFailure();
517 }
518};
519} // namespace
520
522 const LLVMTypeConverter &converter, RewritePatternSet &patterns,
523 bool approximateLog1p, PatternBenefit benefit) {
524 if (approximateLog1p)
525 patterns.add<Log1pOpLowering>(converter, benefit);
526 // clang-format off
527 patterns.add<
528 IsNaNOpLowering,
529 IsFiniteOpLowering,
530 AbsFOpLowering,
531 AbsIOpLowering,
532 CeilOpLowering,
533 CopySignOpLowering,
534 CosOpLowering,
535 CoshOpLowering,
536 AcosOpLowering,
537 CountLeadingZerosOpLowering,
538 CountTrailingZerosOpLowering,
539 CtPopFOpLowering,
540 Exp2OpLowering,
541 ExpM1OpLowering,
542 ExpOpLowering,
543 FPowIOpLowering,
544 FloorOpLowering,
545 FmaOpLowering,
546 ConstrainedFmaOpLowering,
547 Log10OpLowering,
548 Log2OpLowering,
549 LogOpLowering,
550 PowFOpLowering,
551 RoundEvenOpLowering,
552 RoundOpLowering,
553 RsqrtOpLowering,
555 SinOpLowering,
556 SinhOpLowering,
557 ASinOpLowering,
558 SqrtOpLowering,
559 FTruncOpLowering,
560 TanOpLowering,
561 TanhOpLowering,
562 ATanOpLowering,
563 ATan2OpLowering
564 >(converter, benefit);
565 // clang-format on
566}
567
568//===----------------------------------------------------------------------===//
569// ConvertToLLVMPatternInterface implementation
570//===----------------------------------------------------------------------===//
571
572namespace {
573/// Implement the interface to convert Math to LLVM.
574struct MathToLLVMDialectInterface : public ConvertToLLVMPatternInterface {
575 MathToLLVMDialectInterface(Dialect *dialect)
576 : ConvertToLLVMPatternInterface(dialect) {}
577
578 void loadDependentDialects(MLIRContext *context) const final {
579 context->loadDialect<LLVM::LLVMDialect>();
580 }
581
582 /// Hook for derived dialect interface to provide conversion patterns
583 /// and mark dialect legal for the conversion target.
584 void populateConvertToLLVMConversionPatterns(
585 ConversionTarget &target, LLVMTypeConverter &typeConverter,
586 RewritePatternSet &patterns) const final {
587 populateMathToLLVMConversionPatterns(typeConverter, patterns);
588 }
589};
590} // namespace
591
593 registry.addExtension(+[](MLIRContext *ctx, math::MathDialect *dialect) {
594 dialect->addInterfaces<MathToLLVMDialectInterface>();
595 });
596}
return success()
b getContext())
#define add(a, b)
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Definition Pattern.h:239
const LLVMTypeConverter * getTypeConverter() const
Definition Pattern.cpp:29
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Dialects are groups of MLIR operations, types and attributes, as well as behavior associated with the...
Definition Dialect.h:38
Conversion from types to the LLVM IR dialect.
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Basic lowering implementation to rewrite Ops with just one result to the LLVM Dialect.
LogicalResult handleMultidimensionalVectors(Operation *op, ValueRange operands, const LLVMTypeConverter &typeConverter, std::function< Value(Type, ValueRange)> createOperand, ConversionPatternRewriter &rewriter)
bool isCompatibleVectorType(Type type)
Returns true if the given type is a vector type compatible with the LLVM dialect.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
void populateMathToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, bool approximateLog1p=true, PatternBenefit benefit=1)
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
void registerConvertMathToLLVMInterface(DialectRegistry &registry)
LogicalResult matchAndRewrite(math::SincosOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override