22#include "llvm/ADT/FloatingPointMode.h"
25#define GEN_PASS_DEF_CONVERTMATHTOLLVMPASS
26#include "mlir/Conversion/Passes.h.inc"
33template <
typename SourceOp,
typename TargetOp>
36template <
typename SourceOp,
typename TargetOp,
bool FailOnUnsupportedFP = true>
37using ConvertFMFMathToLLVMPattern =
47template <
typename SourceOp,
typename TargetOp,
bool HasRoundingMode,
48 template <
typename,
typename>
typename AttrConvert =
50 bool FailOnUnsupportedFP =
true>
51struct ConstrainedVectorConvertToLLVMPattern
53 FailOnUnsupportedFP> {
54 using VectorConvertToLLVMPattern<
55 SourceOp, TargetOp, AttrConvert,
56 FailOnUnsupportedFP>::VectorConvertToLLVMPattern;
59 matchAndRewrite(SourceOp op,
typename SourceOp::Adaptor adaptor,
60 ConversionPatternRewriter &rewriter)
const override {
61 if (HasRoundingMode !=
static_cast<bool>(op.getRoundingModeAttr()))
63 return VectorConvertToLLVMPattern<
64 SourceOp, TargetOp, AttrConvert,
65 FailOnUnsupportedFP>::matchAndRewrite(op, adaptor, rewriter);
70 ConvertFMFMathToLLVMPattern<math::AbsFOp, LLVM::FAbsOp,
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 =
82using Exp2OpLowering = ConvertFMFMathToLLVMPattern<math::Exp2Op, LLVM::Exp2Op>;
83using ExpOpLowering = ConvertFMFMathToLLVMPattern<math::ExpOp, LLVM::ExpOp>;
84using FloorOpLowering =
85 ConvertFMFMathToLLVMPattern<math::FloorOp, LLVM::FFloorOp>;
87 ConstrainedVectorConvertToLLVMPattern<math::FmaOp, LLVM::FMAOp,
91using ConstrainedFmaOpLowering = ConstrainedVectorConvertToLLVMPattern<
92 math::FmaOp, LLVM::ConstrainedFMAIntr,
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>;
117template <
typename MathOp,
typename LLVMOp>
118struct IntOpWithFlagLowering
120 using ConvertOpToLLVMPattern<
121 MathOp,
true>::ConvertOpToLLVMPattern;
122 using Super = IntOpWithFlagLowering<MathOp, LLVMOp>;
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)
133 auto loc = op.getLoc();
134 auto resultType = op.getResult().getType();
135 auto llvmResultType = typeConverter.convertType(resultType);
139 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
140 rewriter.replaceOpWithNewOp<LLVMOp>(op, llvmResultType,
141 adaptor.getOperand(),
false);
145 if (!isa<VectorType>(resultType))
149 op.getOperation(), adaptor.getOperands(), typeConverter,
150 [&](Type llvm1DVectorTy,
ValueRange operands) {
151 return LLVMOp::create(rewriter, loc, llvm1DVectorTy, operands[0],
158using CountLeadingZerosOpLowering =
159 IntOpWithFlagLowering<math::CountLeadingZerosOp, LLVM::CountLeadingZerosOp>;
160using CountTrailingZerosOpLowering =
161 IntOpWithFlagLowering<math::CountTrailingZerosOp,
162 LLVM::CountTrailingZerosOp>;
163using AbsIOpLowering = IntOpWithFlagLowering<math::AbsIOp, LLVM::AbsOp>;
166struct SincosOpLowering
170 math::SincosOp,
true>::ConvertOpToLLVMPattern;
174 ConversionPatternRewriter &rewriter)
const override {
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)
184 ConvertFastMath<math::SincosOp, LLVM::SincosOp> attrs(op);
186 auto structType = LLVM::LLVMStructType::getLiteral(
187 rewriter.getContext(), {llvmOperandType, llvmOperandType});
189 auto sincosOp = LLVM::SincosOp::create(
191 attrs.getProperties(), attrs.getDiscardableAttrs());
193 auto sinValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 0);
194 auto cosValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 1);
196 rewriter.replaceOp(op, {sinValue, cosValue});
202struct ExpM1OpLowering
205 using ConvertOpToLLVMPattern<
206 math::ExpM1Op,
true>::ConvertOpToLLVMPattern;
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)
217 auto loc = op.getLoc();
218 auto resultType = op.getResult().getType();
219 auto floatType = cast<FloatType>(
221 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
222 ConvertFastMath<math::ExpM1Op, LLVM::ExpOp> expAttrs(op);
223 ConvertFastMath<math::ExpM1Op, LLVM::FSubOp> subAttrs(op);
225 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
226 LLVM::ConstantOp one;
228 one = LLVM::ConstantOp::create(
229 rewriter, loc, llvmOperandType,
230 SplatElementsAttr::get(cast<ShapedType>(llvmOperandType),
234 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
236 auto exp = LLVM::ExpOp::create(rewriter, loc,
TypeRange{llvmOperandType},
238 expAttrs.getProperties(),
239 expAttrs.getDiscardableAttrs());
240 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(
242 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
246 if (!isa<VectorType>(resultType))
247 return rewriter.notifyMatchFailure(op,
"expected vector result type");
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()}),
257 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
259 auto exp = LLVM::ExpOp::create(
261 expAttrs.getProperties(), expAttrs.getDiscardableAttrs());
262 return LLVM::FSubOp::create(
264 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
271struct Log1pOpLowering
274 using ConvertOpToLLVMPattern<
275 math::Log1pOp,
true>::ConvertOpToLLVMPattern;
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");
286 auto loc = op.getLoc();
287 auto resultType = op.getResult().getType();
288 auto floatType = cast<FloatType>(
290 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
291 ConvertFastMath<math::Log1pOp, LLVM::FAddOp> addAttrs(op);
292 ConvertFastMath<math::Log1pOp, LLVM::LogOp> logAttrs(op);
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),
301 : LLVM::ConstantOp::create(rewriter, loc, llvmOperandType,
304 auto add = LLVM::FAddOp::create(rewriter, loc,
TypeRange{llvmOperandType},
306 addAttrs.getProperties(),
307 addAttrs.getDiscardableAttrs());
308 rewriter.replaceOpWithNewOp<LLVM::LogOp>(
310 logAttrs.getProperties(), logAttrs.getDiscardableAttrs());
314 if (!isa<VectorType>(resultType))
315 return rewriter.notifyMatchFailure(op,
"expected vector result type");
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()}),
325 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
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},
333 logAttrs.getDiscardableAttrs());
340struct RsqrtOpLowering
343 using ConvertOpToLLVMPattern<
344 math::RsqrtOp,
true>::ConvertOpToLLVMPattern;
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)
355 auto loc = op.getLoc();
356 auto resultType = op.getResult().getType();
357 auto floatType = cast<FloatType>(
359 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
360 ConvertFastMath<math::RsqrtOp, LLVM::SqrtOp> sqrtAttrs(op);
361 ConvertFastMath<math::RsqrtOp, LLVM::FDivOp> divAttrs(op);
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),
372 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
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>(
380 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
384 if (!isa<VectorType>(resultType))
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()}),
395 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
397 auto sqrt = LLVM::SqrtOp::create(
399 sqrtAttrs.getProperties(), sqrtAttrs.getDiscardableAttrs());
400 return LLVM::FDivOp::create(
402 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
408struct FPowIOpLowering
411 using ConvertOpToLLVMPattern<
412 math::FPowIOp,
true>::ConvertOpToLLVMPattern;
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)
422 auto loc = op.getLoc();
423 Value exponent = adaptor.getRhs();
424 if (isa<VectorType>(op.getRhs().getType())) {
425 SplatElementsAttr splatAttr;
427 return rewriter.notifyMatchFailure(op,
"expected a splat exponent");
429 auto exponentType = typeConverter.convertType(
433 exponent = LLVM::ConstantOp::create(
434 rewriter, loc, exponentType,
435 rewriter.getIntegerAttr(exponentType,
439 ConvertFastMath<math::FPowIOp, LLVM::PowIOp> attrs(op);
441 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
442 rewriter.replaceOpWithNewOp<LLVM::PowIOp>(
444 ValueRange{adaptor.getLhs(), exponent}, attrs.getProperties(),
445 attrs.getDiscardableAttrs());
450 op.getOperation(),
ValueRange{adaptor.getLhs()}, typeConverter,
451 [&](Type llvm1DVectorTy,
ValueRange operands) {
452 return LLVM::PowIOp::create(rewriter, loc, TypeRange{llvm1DVectorTy},
454 attrs.getProperties(),
455 attrs.getDiscardableAttrs());
461struct IsNaNOpLowering
464 using ConvertOpToLLVMPattern<
465 math::IsNaNOp,
true>::ConvertOpToLLVMPattern;
468 matchAndRewrite(math::IsNaNOp op, OpAdaptor adaptor,
469 ConversionPatternRewriter &rewriter)
const override {
470 const auto &typeConverter = *this->getTypeConverter();
472 typeConverter.convertType(adaptor.getOperand().getType());
473 auto resultType = typeConverter.convertType(op.getResult().getType());
474 if (!operandType || !resultType)
477 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
478 op, resultType, adaptor.getOperand(), llvm::fcNan);
483struct IsFiniteOpLowering
486 using ConvertOpToLLVMPattern<
487 math::IsFiniteOp,
true>::ConvertOpToLLVMPattern;
490 matchAndRewrite(math::IsFiniteOp op, OpAdaptor adaptor,
491 ConversionPatternRewriter &rewriter)
const override {
492 const auto &typeConverter = *this->getTypeConverter();
494 typeConverter.convertType(adaptor.getOperand().getType());
495 auto resultType = typeConverter.convertType(op.getResult().getType());
496 if (!operandType || !resultType)
499 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
500 op, resultType, adaptor.getOperand(), llvm::fcFinite);
505struct ConvertMathToLLVMPass
509 void runOnOperation()
override {
514 if (
failed(applyPartialConversion(getOperation(),
target,
515 std::move(patterns))))
524 if (approximateLog1p)
525 patterns.
add<Log1pOpLowering>(converter, benefit);
537 CountLeadingZerosOpLowering,
538 CountTrailingZerosOpLowering,
546 ConstrainedFmaOpLowering,
564 >(converter, benefit);
574struct MathToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
575 MathToLLVMDialectInterface(
Dialect *dialect)
576 : ConvertToLLVMPatternInterface(dialect) {}
578 void loadDependentDialects(MLIRContext *context)
const final {
579 context->loadDialect<LLVM::LLVMDialect>();
584 void populateConvertToLLVMConversionPatterns(
585 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
586 RewritePatternSet &patterns)
const final {
594 dialect->addInterfaces<MathToLLVMDialectInterface>();
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
const LLVMTypeConverter * getTypeConverter() const
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...
Conversion from types to the LLVM IR dialect.
MLIRContext is the top-level object for a collection of MLIR operations.
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.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
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.
void registerConvertMathToLLVMInterface(DialectRegistry ®istry)
LogicalResult matchAndRewrite(math::SincosOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override