diff --git a/passes/techmap/mcm.cc b/passes/techmap/mcm.cc index e7a64f6f6..654f8232e 100644 --- a/passes/techmap/mcm.cc +++ b/passes/techmap/mcm.cc @@ -61,6 +61,7 @@ struct AdderGraph { enum class AdderGraphStatus { success, no_graph, + node_limit, budget_exhausted, }; @@ -111,10 +112,14 @@ struct SearchParams { SearchResult &result; }; +struct SearchState; + class Scorer { public: virtual bool better(const AdderGraphCandidate &candidate, const AdderGraphCandidate &selected, bool direct_target) const = 0; + virtual bool score_intermediate(SearchState &state, const AdderGraphCandidate &candidate, + const pool &remaining, int &score) const = 0; virtual ~Scorer() = default; }; @@ -131,6 +136,10 @@ public: return candidate.depth < selected.depth; return candidate.value < selected.value; } + + bool score_intermediate(SearchState &state, const AdderGraphCandidate &candidate, + const pool &remaining, int &score) const override; + }; struct SearchState { @@ -240,36 +249,12 @@ struct SearchState { return selected.value >= 0; } - // Score an intermediate by how many targets it makes reachable in one step. - bool score_intermediate(const AdderGraphCandidate &candidate, const pool &remaining, int &score) - { - pool unlocked; - for (auto &node : graph.depth) { - if (1 + std::max(candidate.depth, node.second) > config.max_depth) - continue; - for (auto &op : enumerate_pair(candidate.value, node.first)) - if (remaining.count(op.res)) - unlocked.insert(op.res); - if (budget_exhausted) - return false; - } - if (1 + candidate.depth <= config.max_depth) { - for (auto &op : enumerate_pair(candidate.value, candidate.value)) - if (remaining.count(op.res)) - unlocked.insert(op.res); - if (budget_exhausted) - return false; - } - score = GetSize(unlocked); - return true; - } - bool select_intermediate(const CandidateMap &candidates, const pool &remaining, AdderGraphCandidate &selected, const Scorer &scorer) { for (auto &entry : candidates) { AdderGraphCandidate candidate = entry.second; - if (!score_intermediate(candidate, remaining, candidate.score)) + if (!scorer.score_intermediate(*this, candidate, remaining, candidate.score)) return false; if (scorer.better(candidate, selected, false)) selected = candidate; @@ -303,6 +288,31 @@ struct SearchState { } }; +// Score an intermediate by how many targets it makes reachable in one step. +bool HcubScorer::score_intermediate(SearchState &state, const AdderGraphCandidate &candidate, + const pool &remaining, int &score) const +{ + pool unlocked; + for (auto &node : state.graph.depth) { + if (1 + std::max(candidate.depth, node.second) > state.config.max_depth) + continue; + for (auto &op : state.enumerate_pair(candidate.value, node.first)) + if (remaining.count(op.res)) + unlocked.insert(op.res); + if (state.budget_exhausted) + return false; + } + if (1 + candidate.depth <= state.config.max_depth) { + for (auto &op : state.enumerate_pair(candidate.value, candidate.value)) + if (remaining.count(op.res)) + unlocked.insert(op.res); + if (state.budget_exhausted) + return false; + } + score = GetSize(unlocked); + return true; +} + class Search { public: virtual void search(const SearchParams ¶ms, const Scorer &scorer) = 0; @@ -333,7 +343,7 @@ public: if (remaining.empty()) break; if (committed_nodes >= state.config.max_nodes) { - params.result.status = AdderGraphStatus::no_graph; + params.result.status = AdderGraphStatus::node_limit; return; } @@ -624,6 +634,9 @@ struct McmWorker { if (status == AdderGraphStatus::budget_exhausted) log(" mcm: search budget of %lld exhausted for %d constant(s), skipping.\n", config.search.work_budget, GetSize(group.items)); + else if (status == AdderGraphStatus::node_limit) + log(" mcm: maximum node count %d reached for %d constant(s), skipping.\n", + config.search.max_nodes, GetSize(group.items)); else log(" mcm: no adder graph within depth %d for %d constant(s), skipping.\n", config.search.max_depth, GetSize(group.items)); @@ -670,6 +683,9 @@ struct McmPass : public Pass { log(" -max_shift \n"); log(" largest shift considered in an A-operation (default: 12).\n"); log("\n"); + log(" -max_nodes \n"); + log(" maximum nember of nodes (default: 64).\n"); + log("\n"); log(" -min_const \n"); log(" skip constants whose magnitude is below (default: 3).\n"); log("\n"); @@ -705,6 +721,12 @@ struct McmPass : public Pass { log_cmd_error("mcm: -max_shift must be >= 1\n"); continue; } + if (args[argidx] == "-max_nodes" && argidx + 1 < args.size()) { + config.search.max_nodes = atoi(args[++argidx].c_str()); + if (config.search.max_nodes < 1) + log_cmd_error("mcm: -max_nodes must be >= 1\n"); + continue; + } if (args[argidx] == "-min_const" && argidx + 1 < args.size()) { config.min_const = atoi(args[++argidx].c_str()); continue; diff --git a/tests/techmap/mcm.ys b/tests/techmap/mcm.ys index d8ff17ed9..d0bb41619 100644 --- a/tests/techmap/mcm.ys +++ b/tests/techmap/mcm.ys @@ -45,26 +45,40 @@ EOF equiv_opt -assert mcm -force design -load postopt select -assert-count 0 t:$mul -select -assert-count 3 t:$add t:$sub +select -assert-max 3 t:$add t:$sub design -reset -log -header "MCM profitability threshold can reject graphs" +log -header "MCM accepts a graph at the gain threshold" read_verilog <