20#include "llvm/ADT/SmallVectorExtras.h"
23#define GEN_PASS_DEF_CONVERTMATHTOLIBMPASS
24#include "mlir/Conversion/Passes.h.inc"
35 using OpRewritePattern<
Op>::OpRewritePattern;
37 LogicalResult matchAndRewrite(
Op op, PatternRewriter &rewriter)
const final;
43 using OpRewritePattern<
Op>::OpRewritePattern;
45 LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter)
const final;
52 using OpRewritePattern<
Op>::OpRewritePattern;
53 ScalarOpToLibmCall(MLIRContext *context, PatternBenefit benefit,
54 StringRef floatFunc, StringRef doubleFunc)
55 : OpRewritePattern<
Op>(context, benefit), floatFunc(floatFunc),
56 doubleFunc(doubleFunc) {};
58 LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter)
const final;
61 std::string floatFunc, doubleFunc;
64template <
typename OpTy>
67 StringRef doubleFunc) {
68 patterns.
add<VecOpToScalarOp<OpTy>, PromoteOpToF32<OpTy>>(ctx, benefit);
69 patterns.
add<ScalarOpToLibmCall<OpTy>>(ctx, benefit, floatFunc, doubleFunc);
77 auto opType = op.getType();
79 auto vecType = dyn_cast<VectorType>(opType);
83 if (!vecType.hasRank())
85 auto shape = vecType.getShape();
86 int64_t numElements = vecType.getNumElements();
91 FloatAttr::get(vecType.getElementType(), 0.0)));
93 for (
auto linearIndex = 0; linearIndex < numElements; ++linearIndex) {
96 for (
auto input : op->getOperands())
98 vector::ExtractOp::create(rewriter, loc, input, positions));
99 Value scalarOp = Op::create(
100 rewriter, loc,
TypeRange{vecType.getElementType()}, operands,
101 op.
getProperties(), op->getDiscardableAttrDictionary().getValue());
103 vector::InsertOp::create(rewriter, loc, scalarOp,
result, positions);
109template <
typename Op>
112 auto opType = op.getType();
113 if (!isa<Float16Type, BFloat16Type>(opType))
118 auto extendedOperands =
119 llvm::map_to_vector(op->getOperands(), [&](
Value operand) ->
Value {
120 return arith::ExtFOp::create(rewriter, loc, TypeRange{f32},
122 arith::ExtFOp::Properties{});
124 auto newOp = Op::create(rewriter, loc,
TypeRange{f32}, extendedOperands,
126 op->getDiscardableAttrDictionary().getValue());
131template <
typename Op>
133ScalarOpToLibmCall<Op>::matchAndRewrite(
Op op,
135 auto module = SymbolTable::getNearestSymbolTable(op);
136 auto type = op.getType();
137 if (!isa<Float32Type, Float64Type>(type))
140 auto name = type.getIntOrFloatBitWidth() == 64 ? doubleFunc : floatFunc;
141 auto opFunc = dyn_cast_or_null<SymbolOpInterface>(
147 auto opFunctionTy = FunctionType::get(
148 rewriter.
getContext(), op->getOperandTypes(), op->getResultTypes());
149 opFunc = func::FuncOp::create(rewriter, rewriter.
getUnknownLoc(), name,
158 opFunc->setDiscardableAttr(LLVM::LLVMDialect::getReadnoneAttrName(),
173 populatePatternsForOp<math::AbsFOp>(patterns, benefit, ctx,
"fabsf",
"fabs");
174 populatePatternsForOp<math::AcosOp>(patterns, benefit, ctx,
"acosf",
"acos");
175 populatePatternsForOp<math::AcoshOp>(patterns, benefit, ctx,
"acoshf",
177 populatePatternsForOp<math::AsinOp>(patterns, benefit, ctx,
"asinf",
"asin");
178 populatePatternsForOp<math::AsinhOp>(patterns, benefit, ctx,
"asinhf",
180 populatePatternsForOp<math::Atan2Op>(patterns, benefit, ctx,
"atan2f",
182 populatePatternsForOp<math::AtanOp>(patterns, benefit, ctx,
"atanf",
"atan");
183 populatePatternsForOp<math::AtanhOp>(patterns, benefit, ctx,
"atanhf",
185 populatePatternsForOp<math::CbrtOp>(patterns, benefit, ctx,
"cbrtf",
"cbrt");
186 populatePatternsForOp<math::CeilOp>(patterns, benefit, ctx,
"ceilf",
"ceil");
187 populatePatternsForOp<math::CosOp>(patterns, benefit, ctx,
"cosf",
"cos");
188 populatePatternsForOp<math::CoshOp>(patterns, benefit, ctx,
"coshf",
"cosh");
189 populatePatternsForOp<math::ErfOp>(patterns, benefit, ctx,
"erff",
"erf");
190 populatePatternsForOp<math::ErfcOp>(patterns, benefit, ctx,
"erfcf",
"erfc");
191 populatePatternsForOp<math::ExpOp>(patterns, benefit, ctx,
"expf",
"exp");
192 populatePatternsForOp<math::Exp2Op>(patterns, benefit, ctx,
"exp2f",
"exp2");
193 populatePatternsForOp<math::ExpM1Op>(patterns, benefit, ctx,
"expm1f",
195 populatePatternsForOp<math::FloorOp>(patterns, benefit, ctx,
"floorf",
197 populatePatternsForOp<math::FmaOp>(patterns, benefit, ctx,
"fmaf",
"fma");
198 populatePatternsForOp<math::LogOp>(patterns, benefit, ctx,
"logf",
"log");
199 populatePatternsForOp<math::Log2Op>(patterns, benefit, ctx,
"log2f",
"log2");
200 populatePatternsForOp<math::Log10Op>(patterns, benefit, ctx,
"log10f",
202 populatePatternsForOp<math::Log1pOp>(patterns, benefit, ctx,
"log1pf",
204 populatePatternsForOp<math::PowFOp>(patterns, benefit, ctx,
"powf",
"pow");
205 populatePatternsForOp<math::RoundEvenOp>(patterns, benefit, ctx,
"roundevenf",
207 populatePatternsForOp<math::RoundOp>(patterns, benefit, ctx,
"roundf",
209 populatePatternsForOp<math::SinOp>(patterns, benefit, ctx,
"sinf",
"sin");
210 populatePatternsForOp<math::SinhOp>(patterns, benefit, ctx,
"sinhf",
"sinh");
211 populatePatternsForOp<math::SqrtOp>(patterns, benefit, ctx,
"sqrtf",
"sqrt");
212 populatePatternsForOp<math::RsqrtOp>(patterns, benefit, ctx,
"rsqrtf",
214 populatePatternsForOp<math::TanOp>(patterns, benefit, ctx,
"tanf",
"tan");
215 populatePatternsForOp<math::TanhOp>(patterns, benefit, ctx,
"tanhf",
"tanh");
216 populatePatternsForOp<math::TruncOp>(patterns, benefit, ctx,
"truncf",
221struct ConvertMathToLibmPass
222 :
public impl::ConvertMathToLibmPassBase<ConvertMathToLibmPass> {
223 void runOnOperation()
override;
227void ConvertMathToLibmPass::runOnOperation() {
228 auto module = getOperation();
234 target.addLegalDialect<arith::ArithDialect, BuiltinDialect, func::FuncDialect,
235 vector::VectorDialect>();
236 target.addIllegalDialect<math::MathDialect>();
237 if (
failed(applyPartialConversion(module,
target, std::move(patterns))))
MLIRContext * getContext() const
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
MLIRContext is the top-level object for a collection of MLIR operations.
RAII guard to reset the insertion point of the builder when destroyed.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Location getLoc()
The source location the operation was defined or derived from.
This provides public APIs that all operations should have.
InferredProperties< T > & getProperties()
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
static Operation * lookupSymbolIn(Operation *op, StringAttr symbol)
Returns the operation registered with the given symbol name with the regions of 'symbolTableOp'.
This class provides an abstraction over the various different ranges of value types.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
NestedPattern Op(FilterFunctionType filter=defaultFilterFunction)
Include the generated interface declarations.
void populateMathToLibmConversionPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate the given list with patterns that convert from Math to Libm calls.
SmallVector< int64_t > computeStrides(ArrayRef< int64_t > sizes)
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...