#include "mlir/Dialect/Math/IR/Math.h"
#include "mlir/Dialect/Math/Transforms/Passes.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/Pass/Pass.h"
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
using namespace mlir;
namespace {
struct TestMathAlgebraicSimplificationPass
: public PassWrapper<TestMathAlgebraicSimplificationPass, OperationPass<>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(
TestMathAlgebraicSimplificationPass)
void runOnOperation() override;
void getDependentDialects(DialectRegistry ®istry) const override {
registry.insert<vector::VectorDialect, math::MathDialect>();
}
StringRef getArgument() const final {
return "test-math-algebraic-simplification";
}
StringRef getDescription() const final {
return "Test math algebraic simplification";
}
};
}
void TestMathAlgebraicSimplificationPass::runOnOperation() {
RewritePatternSet patterns(&getContext());
populateMathAlgebraicSimplificationPatterns(patterns);
(void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns));
}
namespace mlir {
namespace test {
void registerTestMathAlgebraicSimplificationPass() {
PassRegistration<TestMathAlgebraicSimplificationPass>();
}
}
}