20struct MathToEmitCDialectInterface :
public ConvertToEmitCPatternInterface {
21 MathToEmitCDialectInterface(Dialect *dialect)
22 : ConvertToEmitCPatternInterface(dialect) {}
26 void populateConvertToEmitCConversionPatterns(
27 ConversionTarget &
target, TypeConverter &typeConverter,
28 RewritePatternSet &patterns, std::optional<bool> lowerToCpp)
const final {
30 patterns, lowerToCpp.value_or(
true) ? emitc::LanguageTarget::cpp11
31 : emitc::LanguageTarget::c99);
35template <
typename OpType>
37 std::string calleeStr;
41 LowerToEmitCCallOpaque(MLIRContext *context, std::string calleeStr,
43 : OpRewritePattern<OpType>(context), calleeStr(std::move(calleeStr)),
44 languageTarget(languageTarget) {}
46 LogicalResult matchAndRewrite(OpType op,
47 PatternRewriter &rewriter)
const override;
50template <
typename OpType>
51LogicalResult LowerToEmitCCallOpaque<OpType>::matchAndRewrite(
53 if (!llvm::all_of(op->getOperandTypes(),
54 llvm::IsaPred<Float32Type, Float64Type>) ||
55 !llvm::all_of(op->getResultTypes(),
56 llvm::IsaPred<Float32Type, Float64Type>))
59 "expected all operands and results to be of type f32 or f64");
60 std::string modifiedCalleeStr = calleeStr;
62 modifiedCalleeStr =
"std::" + calleeStr;
64 auto operandType = op->getOperandTypes()[0];
65 if (operandType.isF32())
66 modifiedCalleeStr = calleeStr +
"f";
69 op, op.getType(), modifiedCalleeStr, op->getOperands());
77 dialect->addInterfaces<MathToEmitCDialectInterface>();
86 patterns.
insert<LowerToEmitCCallOpaque<math::FloorOp>>(context,
"floor",
88 patterns.
insert<LowerToEmitCCallOpaque<math::RoundOp>>(context,
"round",
90 patterns.
insert<LowerToEmitCCallOpaque<math::RoundEvenOp>>(
91 context,
"roundeven", languageTarget);
92 patterns.
insert<LowerToEmitCCallOpaque<math::ExpOp>>(context,
"exp",
94 patterns.
insert<LowerToEmitCCallOpaque<math::CosOp>>(context,
"cos",
96 patterns.
insert<LowerToEmitCCallOpaque<math::SinOp>>(context,
"sin",
98 patterns.
insert<LowerToEmitCCallOpaque<math::AcosOp>>(context,
"acos",
100 patterns.
insert<LowerToEmitCCallOpaque<math::AsinOp>>(context,
"asin",
102 patterns.
insert<LowerToEmitCCallOpaque<math::Atan2Op>>(context,
"atan2",
104 patterns.
insert<LowerToEmitCCallOpaque<math::CeilOp>>(context,
"ceil",
106 patterns.
insert<LowerToEmitCCallOpaque<math::AbsFOp>>(context,
"fabs",
108 patterns.
insert<LowerToEmitCCallOpaque<math::PowFOp>>(context,
"pow",
110 patterns.
insert<LowerToEmitCCallOpaque<math::SqrtOp>>(context,
"sqrt",
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.
MLIRContext is the top-level object for a collection of MLIR operations.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
MLIRContext * getContext() const
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
LanguageTarget
Enum to specify the language target for EmitC code generation.
Include the generated interface declarations.
void registerConvertMathToEmitCInterface(DialectRegistry ®istry)
void populateConvertMathToEmitCPatterns(RewritePatternSet &patterns, emitc::LanguageTarget languageTarget)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...