21#include "llvm/ADT/FloatingPointMode.h"
24#define GEN_PASS_DEF_CONVERTMATHTOLLVMPASS
25#include "mlir/Conversion/Passes.h.inc"
32template <
typename SourceOp,
typename TargetOp>
35template <
typename SourceOp,
typename TargetOp,
bool FailOnUnsupportedFP = true>
36using ConvertFMFMathToLLVMPattern =
46template <
typename SourceOp,
typename TargetOp,
bool HasRoundingMode,
47 template <
typename,
typename>
typename AttrConvert =
49 bool FailOnUnsupportedFP =
true>
50struct ConstrainedVectorConvertToLLVMPattern
52 FailOnUnsupportedFP> {
53 using VectorConvertToLLVMPattern<
54 SourceOp, TargetOp, AttrConvert,
55 FailOnUnsupportedFP>::VectorConvertToLLVMPattern;
58 matchAndRewrite(SourceOp op,
typename SourceOp::Adaptor adaptor,
59 ConversionPatternRewriter &rewriter)
const override {
60 if (HasRoundingMode !=
static_cast<bool>(op.getRoundingModeAttr()))
62 return VectorConvertToLLVMPattern<
63 SourceOp, TargetOp, AttrConvert,
64 FailOnUnsupportedFP>::matchAndRewrite(op, adaptor, rewriter);
69 ConvertFMFMathToLLVMPattern<math::AbsFOp, LLVM::FAbsOp,
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 =
81using Exp2OpLowering = ConvertFMFMathToLLVMPattern<math::Exp2Op, LLVM::Exp2Op>;
82using ExpOpLowering = ConvertFMFMathToLLVMPattern<math::ExpOp, LLVM::ExpOp>;
83using FloorOpLowering =
84 ConvertFMFMathToLLVMPattern<math::FloorOp, LLVM::FFloorOp>;
86 ConstrainedVectorConvertToLLVMPattern<math::FmaOp, LLVM::FMAOp,
90using ConstrainedFmaOpLowering = ConstrainedVectorConvertToLLVMPattern<
91 math::FmaOp, LLVM::ConstrainedFMAIntr,
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>;
118template <
typename MathOp,
typename LLVMOp>
119struct IntOpWithFlagLowering
121 using ConvertOpToLLVMPattern<
122 MathOp,
true>::ConvertOpToLLVMPattern;
123 using Super = IntOpWithFlagLowering<MathOp, LLVMOp>;
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)
134 auto loc = op.getLoc();
135 auto resultType = op.getResult().getType();
136 auto llvmResultType = typeConverter.convertType(resultType);
140 if (!isa<LLVM::LLVMArrayType>(llvmOperandType)) {
141 rewriter.replaceOpWithNewOp<LLVMOp>(op, llvmResultType,
142 adaptor.getOperand(),
false);
146 if (!isa<VectorType>(resultType))
150 op.getOperation(), adaptor.getOperands(), typeConverter,
151 [&](Type llvm1DVectorTy,
ValueRange operands) {
152 return LLVMOp::create(rewriter, loc, llvm1DVectorTy, operands[0],
159using CountLeadingZerosOpLowering =
160 IntOpWithFlagLowering<math::CountLeadingZerosOp, LLVM::CountLeadingZerosOp>;
161using CountTrailingZerosOpLowering =
162 IntOpWithFlagLowering<math::CountTrailingZerosOp,
163 LLVM::CountTrailingZerosOp>;
164using AbsIOpLowering = IntOpWithFlagLowering<math::AbsIOp, LLVM::AbsOp>;
167struct SincosOpLowering
171 math::SincosOp,
true>::ConvertOpToLLVMPattern;
175 ConversionPatternRewriter &rewriter)
const override {
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)
185 ConvertFastMath<math::SincosOp, LLVM::SincosOp> attrs(op);
187 auto structType = LLVM::LLVMStructType::getLiteral(
188 rewriter.getContext(), {llvmOperandType, llvmOperandType});
190 auto sincosOp = LLVM::SincosOp::create(
192 attrs.getProperties(), attrs.getDiscardableAttrs());
194 auto sinValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 0);
195 auto cosValue = LLVM::ExtractValueOp::create(rewriter, loc, sincosOp, 1);
197 rewriter.replaceOp(op, {sinValue, cosValue});
203struct ExpM1OpLowering
206 using ConvertOpToLLVMPattern<
207 math::ExpM1Op,
true>::ConvertOpToLLVMPattern;
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)
218 auto loc = op.getLoc();
219 auto resultType = op.getResult().getType();
220 auto floatType = cast<FloatType>(
222 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
223 ConvertFastMath<math::ExpM1Op, LLVM::ExpOp> expAttrs(op);
224 ConvertFastMath<math::ExpM1Op, LLVM::FSubOp> subAttrs(op);
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),
235 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
237 auto exp = LLVM::ExpOp::create(rewriter, loc,
TypeRange{llvmOperandType},
239 expAttrs.getProperties(),
240 expAttrs.getDiscardableAttrs());
241 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(
243 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
247 if (!isa<VectorType>(resultType))
248 return rewriter.notifyMatchFailure(op,
"expected vector result type");
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()}),
258 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
260 auto exp = LLVM::ExpOp::create(
262 expAttrs.getProperties(), expAttrs.getDiscardableAttrs());
263 return LLVM::FSubOp::create(
265 subAttrs.getProperties(), subAttrs.getDiscardableAttrs());
272struct Log1pOpLowering
275 using ConvertOpToLLVMPattern<
276 math::Log1pOp,
true>::ConvertOpToLLVMPattern;
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");
287 auto loc = op.getLoc();
288 auto resultType = op.getResult().getType();
289 auto floatType = cast<FloatType>(
291 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
292 ConvertFastMath<math::Log1pOp, LLVM::FAddOp> addAttrs(op);
293 ConvertFastMath<math::Log1pOp, LLVM::LogOp> logAttrs(op);
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),
302 : LLVM::ConstantOp::create(rewriter, loc, llvmOperandType,
305 auto add = LLVM::FAddOp::create(rewriter, loc,
TypeRange{llvmOperandType},
307 addAttrs.getProperties(),
308 addAttrs.getDiscardableAttrs());
309 rewriter.replaceOpWithNewOp<LLVM::LogOp>(
311 logAttrs.getProperties(), logAttrs.getDiscardableAttrs());
315 if (!isa<VectorType>(resultType))
316 return rewriter.notifyMatchFailure(op,
"expected vector result type");
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()}),
326 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
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},
334 logAttrs.getDiscardableAttrs());
341struct RsqrtOpLowering
344 using ConvertOpToLLVMPattern<
345 math::RsqrtOp,
true>::ConvertOpToLLVMPattern;
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)
356 auto loc = op.getLoc();
357 auto resultType = op.getResult().getType();
358 auto floatType = cast<FloatType>(
360 auto floatOne = rewriter.getFloatAttr(floatType, 1.0);
361 ConvertFastMath<math::RsqrtOp, LLVM::SqrtOp> sqrtAttrs(op);
362 ConvertFastMath<math::RsqrtOp, LLVM::FDivOp> divAttrs(op);
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),
373 LLVM::ConstantOp::create(rewriter, loc, llvmOperandType, floatOne);
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>(
381 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
385 if (!isa<VectorType>(resultType))
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()}),
396 auto one = LLVM::ConstantOp::create(rewriter, loc, llvm1DVectorTy,
398 auto sqrt = LLVM::SqrtOp::create(
400 sqrtAttrs.getProperties(), sqrtAttrs.getDiscardableAttrs());
401 return LLVM::FDivOp::create(
403 divAttrs.getProperties(), divAttrs.getDiscardableAttrs());
409struct IsNaNOpLowering
412 using ConvertOpToLLVMPattern<
413 math::IsNaNOp,
true>::ConvertOpToLLVMPattern;
416 matchAndRewrite(math::IsNaNOp op, OpAdaptor adaptor,
417 ConversionPatternRewriter &rewriter)
const override {
418 const auto &typeConverter = *this->getTypeConverter();
420 typeConverter.convertType(adaptor.getOperand().getType());
421 auto resultType = typeConverter.convertType(op.getResult().getType());
422 if (!operandType || !resultType)
425 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
426 op, resultType, adaptor.getOperand(), llvm::fcNan);
431struct IsFiniteOpLowering
434 using ConvertOpToLLVMPattern<
435 math::IsFiniteOp,
true>::ConvertOpToLLVMPattern;
438 matchAndRewrite(math::IsFiniteOp op, OpAdaptor adaptor,
439 ConversionPatternRewriter &rewriter)
const override {
440 const auto &typeConverter = *this->getTypeConverter();
442 typeConverter.convertType(adaptor.getOperand().getType());
443 auto resultType = typeConverter.convertType(op.getResult().getType());
444 if (!operandType || !resultType)
447 rewriter.replaceOpWithNewOp<LLVM::IsFPClass>(
448 op, resultType, adaptor.getOperand(), llvm::fcFinite);
453struct ConvertMathToLLVMPass
454 :
public impl::ConvertMathToLLVMPassBase<ConvertMathToLLVMPass> {
457 void runOnOperation()
override {
462 if (
failed(applyPartialConversion(getOperation(),
target,
463 std::move(patterns))))
472 if (approximateLog1p)
473 patterns.
add<Log1pOpLowering>(converter, benefit);
485 CountLeadingZerosOpLowering,
486 CountTrailingZerosOpLowering,
494 ConstrainedFmaOpLowering,
512 >(converter, benefit);
522struct MathToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
523 MathToLLVMDialectInterface(
Dialect *dialect)
524 : ConvertToLLVMPatternInterface(dialect) {}
526 void loadDependentDialects(MLIRContext *context)
const final {
527 context->loadDialect<LLVM::LLVMDialect>();
532 void populateConvertToLLVMConversionPatterns(
533 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
534 RewritePatternSet &patterns)
const final {
542 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
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)
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 ®istry)
LogicalResult matchAndRewrite(math::SincosOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override