19#define GEN_PASS_DEF_COMPOSITEFIXEDPOINTPASS
20#include "mlir/Transforms/Passes.h.inc"
26struct CompositeFixedPointPass final
27 :
public impl::CompositeFixedPointPassBase<CompositeFixedPointPass> {
28 using CompositeFixedPointPassBase::CompositeFixedPointPassBase;
30 CompositeFixedPointPass(
31 std::string name_, llvm::function_ref<
void(OpPassManager &)> populateFunc,
33 name = std::move(name_);
34 maxIter = maxIterations;
35 convergenceFailureAction = convergenceFailureActionArg;
36 populateFunc(dynamicPM);
38 llvm::raw_string_ostream os(pipelineStr);
41 [&]() { os <<
","; });
44 LogicalResult initializeOptions(
46 function_ref<LogicalResult(
const Twine &)> errorHandler)
override {
47 if (
failed(CompositeFixedPointPassBase::initializeOptions(
options,
52 return errorHandler(
"Failed to parse composite pass pipeline");
57 LogicalResult
initialize(MLIRContext *context)
override {
59 return emitError(UnknownLoc::get(context))
60 <<
"Invalid maxIterations value: " << maxIter <<
"\n";
65 void getDependentDialects(DialectRegistry ®istry)
const override {
66 dynamicPM.getDependentDialects(registry);
69 void runOnOperation()
override {
70 auto *op = getOperation();
71 OperationFingerPrint fp(op);
74 int maxIterVal = maxIter;
76 if (
failed(runPipeline(dynamicPM, op)))
77 return signalPassFailure();
79 if (currentIter++ >= maxIterVal) {
80 std::string message = (
"Composite pass \"" + llvm::Twine(name) +
81 "\"+ didn't converge in " +
82 llvm::Twine(maxIterVal) +
" iterations")
84 switch (convergenceFailureAction) {
85 case ConvergenceFailureAction::Warn:
86 op->emitWarning(message);
88 case ConvergenceFailureAction::Error:
89 op->emitError(message);
90 return signalPassFailure();
91 case ConvergenceFailureAction::Silent:
97 OperationFingerPrint newFp(op);
106 llvm::StringRef getName()
const override {
return name; }
109 OpPassManager dynamicPM;
117 return std::make_unique<CompositeFixedPointPass>(
118 std::move(name), populateFunc, maxIterations, convergenceFailureAction);
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
static llvm::ManagedStatic< PassManagerOptions > options
This class represents a pass manager that runs passes on either a specific operation type,...
void printAsTextualPipeline(raw_ostream &os, bool pretty=false)
Prints out the pass in the textual representation of pipelines.
Include the generated interface declarations.
std::unique_ptr< Pass > createCompositeFixedPointPass(std::string name, llvm::function_ref< void(OpPassManager &)> populateFunc, int maxIterations=10, ConvergenceFailureAction convergenceFailureAction=ConvergenceFailureAction::Warn)
Create composite pass, which runs provided set of passes until fixed point or maximum number of itera...
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
ConvergenceFailureAction
Action to take when CompositeFixedPointPass fails to converge within its configured maximum number of...
LogicalResult parsePassPipeline(StringRef pipeline, OpPassManager &pm, raw_ostream &errorStream=llvm::errs())
Parse the textual representation of a pass pipeline, adding the result to 'pm' on success.
llvm::function_ref< Fn > function_ref