refactor with scorer and search classes

This commit is contained in:
Tharaka
2026-10-02 09:26:57 +02:00
committed by nella
parent a2c891184e
commit 489d90f893
+82 -39
View File
@@ -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<int> &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::pair<int, int>, std::vector<AOp>> 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<int> &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<int> &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<int> &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<int> &targets_in)
class Search {
public:
virtual void search(const SearchParams &params, const Scorer &scorer) = 0;
virtual ~Search() = default;
};
class HcubSearch : public Search {
public:
void search(const SearchParams &params, const Scorer &scorer) override
{
SearchState state(params.config);
params.result = SearchResult();
pool<int> 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<int> remaining = find_remaining_targets(targets);
pool<int> 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++;