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