MLIR 24.0.0git
MathToEmitC.cpp
Go to the documentation of this file.
1//===- MathToEmitC.cpp - Math to EmitC Patterns -----------------*- C++ -*-===//
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
15
16using namespace mlir;
17
18namespace {
19/// Implement the interface to convert Math to EmitC.
20struct MathToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
21 MathToEmitCDialectInterface(Dialect *dialect)
22 : ConvertToEmitCPatternInterface(dialect) {}
23
24 /// Hook for derived dialect interface to provide conversion patterns
25 /// and mark dialect legal for the conversion target.
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);
32 }
33};
34
35template <typename OpType>
36class LowerToEmitCCallOpaque : public OpRewritePattern<OpType> {
37 std::string calleeStr;
38 emitc::LanguageTarget languageTarget;
39
40public:
41 LowerToEmitCCallOpaque(MLIRContext *context, std::string calleeStr,
42 emitc::LanguageTarget languageTarget)
43 : OpRewritePattern<OpType>(context), calleeStr(std::move(calleeStr)),
44 languageTarget(languageTarget) {}
45
46 LogicalResult matchAndRewrite(OpType op,
47 PatternRewriter &rewriter) const override;
48};
49
50template <typename OpType>
51LogicalResult LowerToEmitCCallOpaque<OpType>::matchAndRewrite(
52 OpType op, PatternRewriter &rewriter) const {
53 if (!llvm::all_of(op->getOperandTypes(),
54 llvm::IsaPred<Float32Type, Float64Type>) ||
55 !llvm::all_of(op->getResultTypes(),
56 llvm::IsaPred<Float32Type, Float64Type>))
57 return rewriter.notifyMatchFailure(
58 op.getLoc(),
59 "expected all operands and results to be of type f32 or f64");
60 std::string modifiedCalleeStr = calleeStr;
61 if (languageTarget == emitc::LanguageTarget::cpp11) {
62 modifiedCalleeStr = "std::" + calleeStr;
63 } else if (languageTarget == emitc::LanguageTarget::c99) {
64 auto operandType = op->getOperandTypes()[0];
65 if (operandType.isF32())
66 modifiedCalleeStr = calleeStr + "f";
67 }
68 rewriter.replaceOpWithNewOp<emitc::CallOpaqueOp>(
69 op, op.getType(), modifiedCalleeStr, op->getOperands());
70 return success();
71}
72
73} // namespace
74
76 registry.addExtension(+[](MLIRContext *ctx, math::MathDialect *dialect) {
77 dialect->addInterfaces<MathToEmitCDialectInterface>();
78 });
79}
80
81// Populates patterns to replace `math` operations with `emitc.call_opaque`,
82// using function names consistent with those in <math.h>.
84 RewritePatternSet &patterns, emitc::LanguageTarget languageTarget) {
85 auto *context = patterns.getContext();
86 patterns.insert<LowerToEmitCCallOpaque<math::FloorOp>>(context, "floor",
87 languageTarget);
88 patterns.insert<LowerToEmitCCallOpaque<math::RoundOp>>(context, "round",
89 languageTarget);
90 patterns.insert<LowerToEmitCCallOpaque<math::RoundEvenOp>>(
91 context, "roundeven", languageTarget);
92 patterns.insert<LowerToEmitCCallOpaque<math::ExpOp>>(context, "exp",
93 languageTarget);
94 patterns.insert<LowerToEmitCCallOpaque<math::CosOp>>(context, "cos",
95 languageTarget);
96 patterns.insert<LowerToEmitCCallOpaque<math::SinOp>>(context, "sin",
97 languageTarget);
98 patterns.insert<LowerToEmitCCallOpaque<math::AcosOp>>(context, "acos",
99 languageTarget);
100 patterns.insert<LowerToEmitCCallOpaque<math::AsinOp>>(context, "asin",
101 languageTarget);
102 patterns.insert<LowerToEmitCCallOpaque<math::Atan2Op>>(context, "atan2",
103 languageTarget);
104 patterns.insert<LowerToEmitCCallOpaque<math::CeilOp>>(context, "ceil",
105 languageTarget);
106 patterns.insert<LowerToEmitCCallOpaque<math::AbsFOp>>(context, "fabs",
107 languageTarget);
108 patterns.insert<LowerToEmitCCallOpaque<math::PowFOp>>(context, "pow",
109 languageTarget);
110 patterns.insert<LowerToEmitCCallOpaque<math::SqrtOp>>(context, "sqrt",
111 languageTarget);
112}
return success()
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.
Definition MLIRContext.h:63
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.
Definition MathToEmitC.h:17
Include the generated interface declarations.
void registerConvertMathToEmitCInterface(DialectRegistry &registry)
void populateConvertMathToEmitCPatterns(RewritePatternSet &patterns, emitc::LanguageTarget languageTarget)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...