From 489d90f89380ea318bb9a1f9e2dc922eabbd15b8 Mon Sep 17 00:00:00 2001 From: Tharaka Date: Tue, 29 Sep 2026 14:52:59 +0300 Subject: [PATCH] refactor with scorer and search classes --- passes/techmap/mcm.cc | 121 ++++++++++++++++++++++++++++-------------- 1 file changed, 82 insertions(+), 39 deletions(-) diff --git a/passes/techmap/mcm.cc b/passes/techmap/mcm.cc index d24eb028a..e7a64f6f6 100644 --- a/passes/techmap/mcm.cc +++ b/passes/techmap/mcm.cc @@ -100,7 +100,40 @@ static int naf_weight(int n) return weight; } -struct AdderGraphSolver { +struct SearchResult { + AdderGraph graph; + AdderGraphStatus status = AdderGraphStatus::no_graph; +}; + +struct SearchParams { + const AdderGraphConfig &config; + const std::vector &targets; + SearchResult &result; +}; + +class Scorer { +public: + virtual bool better(const AdderGraphCandidate &candidate, const AdderGraphCandidate &selected, + bool direct_target) const = 0; + virtual ~Scorer() = default; +}; + +class HcubScorer : public Scorer { +public: + bool better(const AdderGraphCandidate &candidate, const AdderGraphCandidate &selected, + bool direct_target) const override + { + if (selected.value < 0) + return true; + if (!direct_target && candidate.score != selected.score) + return candidate.score > selected.score; + if (candidate.depth != selected.depth) + return candidate.depth < selected.depth; + return candidate.value < selected.value; + } +}; + +struct SearchState { AdderGraphConfig config; int limit; long long search_work; @@ -108,7 +141,7 @@ struct AdderGraphSolver { AdderGraph graph; dict, std::vector> op_cache; - explicit AdderGraphSolver(const AdderGraphConfig &config) + explicit SearchState(const AdderGraphConfig &config) : config(config), limit(1), search_work(0), budget_exhausted(false) {} @@ -196,18 +229,13 @@ struct AdderGraphSolver { // Prefer a target we can produce immediately. bool select_direct_target(const CandidateMap &candidates, const pool &remaining, - AdderGraphCandidate &selected) const + AdderGraphCandidate &selected, const Scorer &scorer) const { for (auto &entry : candidates) { if (!remaining.count(entry.first)) continue; - int depth = entry.second.depth; - if (selected.value < 0 || depth < selected.depth || - (depth == selected.depth && entry.first < selected.value)) { - selected.value = entry.first; - selected.depth = depth; - selected.op = entry.second.op; - } + if (scorer.better(entry.second, selected, true)) + selected = entry.second; } return selected.value >= 0; } @@ -237,28 +265,26 @@ struct AdderGraphSolver { } bool select_intermediate(const CandidateMap &candidates, const pool &remaining, - AdderGraphCandidate &selected) + AdderGraphCandidate &selected, const Scorer &scorer) { for (auto &entry : candidates) { AdderGraphCandidate candidate = entry.second; if (!score_intermediate(candidate, remaining, candidate.score)) return false; - if (selected.value < 0 || candidate.score > selected.score || - (candidate.score == selected.score && (candidate.depth < selected.depth || - (candidate.depth == selected.depth && candidate.value < selected.value)))) + if (scorer.better(candidate, selected, false)) selected = candidate; } return selected.value >= 0; } bool select_next_candidate(const CandidateMap &candidates, const pool &remaining, - AdderGraphCandidate &selected) + AdderGraphCandidate &selected, const Scorer &scorer) { - if (select_direct_target(candidates, remaining, selected)) + if (select_direct_target(candidates, remaining, selected, scorer)) return true; if (candidates.empty()) return false; - return select_intermediate(candidates, remaining, selected); + return select_intermediate(candidates, remaining, selected, scorer); } void commit(const AdderGraphCandidate &candidate) @@ -275,45 +301,58 @@ struct AdderGraphSolver { remaining.insert(target); return remaining; } +}; - AdderGraphStatus solve(const std::vector &targets_in) +class Search { +public: + virtual void search(const SearchParams ¶ms, const Scorer &scorer) = 0; + virtual ~Search() = default; +}; + +class HcubSearch : public Search { +public: + void search(const SearchParams ¶ms, const Scorer &scorer) override { + SearchState state(params.config); + params.result = SearchResult(); pool targets; - for (int target : targets_in) { + for (int target : params.targets) { int odd_target = oddify(target); if (odd_target > 1) targets.insert(odd_target); } - limit = 1; + state.limit = 1; for (int target : targets) - limit = std::max(limit, target); + state.limit = std::max(state.limit, target); - graph.depth.clear(); - graph.ops.clear(); - op_cache.clear(); - search_work = 0; - budget_exhausted = false; - graph.depth[1] = 0; + state.graph.depth[1] = 0; int committed_nodes = 0; while (true) { - pool remaining = find_remaining_targets(targets); + pool remaining = state.find_remaining_targets(targets); if (remaining.empty()) break; - if (committed_nodes >= config.max_nodes) - return AdderGraphStatus::no_graph; + if (committed_nodes >= state.config.max_nodes) { + params.result.status = AdderGraphStatus::no_graph; + return; + } CandidateMap candidates; - if (!build_frontier(candidates)) - return budget_exhausted ? AdderGraphStatus::budget_exhausted : AdderGraphStatus::no_graph; + if (!state.build_frontier(candidates)) { + params.result.status = state.budget_exhausted ? AdderGraphStatus::budget_exhausted : AdderGraphStatus::no_graph; + return; + } AdderGraphCandidate selected; - if (!select_next_candidate(candidates, remaining, selected)) - return budget_exhausted ? AdderGraphStatus::budget_exhausted : AdderGraphStatus::no_graph; - commit(selected); + if (!state.select_next_candidate(candidates, remaining, selected, scorer)) { + params.result.status = state.budget_exhausted ? AdderGraphStatus::budget_exhausted : AdderGraphStatus::no_graph; + return; + } + state.commit(selected); committed_nodes++; } - return AdderGraphStatus::success; + params.result.graph = std::move(state.graph); + params.result.status = AdderGraphStatus::success; } }; @@ -575,8 +614,12 @@ struct McmWorker { for (auto &item : group.items) targets.push_back(item.coefficient); - AdderGraphSolver solver(config.search); - AdderGraphStatus status = solver.solve(targets); + HcubSearch hcub_search; + Search &search = hcub_search; + HcubScorer scorer; + SearchResult result; + search.search(SearchParams{config.search, targets, result}, scorer); + AdderGraphStatus status = result.status; if (status != AdderGraphStatus::success) { if (status == AdderGraphStatus::budget_exhausted) log(" mcm: search budget of %lld exhausted for %d constant(s), skipping.\n", @@ -588,7 +631,7 @@ struct McmWorker { } McmPlan plan; - if (!prepare_plan(group, solver.graph, plan)) + if (!prepare_plan(group, result.graph, plan)) return; emit_plan(group, plan); n_groups++;