mcm: refactor and move score_intermediate into scorer

This commit is contained in:
Tharaka
2026-10-02 09:26:57 +02:00
committed by nella
parent 489d90f893
commit ebaecb50bf
2 changed files with 70 additions and 34 deletions
+48 -26
View File
@@ -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<int> &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<int> &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<int> &remaining, int &score)
{
pool<int> 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<int> &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<int> &remaining, int &score) const
{
pool<int> 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 &params, 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>\n");
log(" largest shift considered in an A-operation (default: 12).\n");
log("\n");
log(" -max_nodes <n>\n");
log(" maximum nember of nodes (default: 64).\n");
log("\n");
log(" -min_const <n>\n");
log(" skip constants whose magnitude is below <n> (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;
+22 -8
View File
@@ -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 <<EOF
module top(
input wire [7:0] x,
input wire [7:0] singleton_x,
output wire [15:0] singleton,
output wire [15:0] a,
output wire [15:0] b
);
assign singleton = singleton_x * 16'd397;
assign a = x * 16'd3;
assign b = x * 16'd5;
assign b = x * 16'd21;
endmodule
EOF
mcm -min_gain 100
select -assert-count 3 t:$mul
equiv_opt -assert mcm -min_gain 33
design -load postopt
select -assert-count 0 t:$mul
design -reset
log -header "MCM rejects a graph above its gain"
read_verilog <<EOF
module top(
input wire [7:0] x,
output wire [15:0] a,
output wire [15:0] b
);
assign a = x * 16'd3;
assign b = x * 16'd21;
endmodule
EOF
mcm -min_gain 34
select -assert-count 2 t:$mul
design -reset