use crate::cnf::{Clause, CnfFormula};
use crate::error::VitriError;
use crate::vtree::{VarId, Vtree, VtreeIdx};
use std::collections::{HashMap, VecDeque};
pub(crate) mod agg;
pub(crate) mod tables;
fn covered_by(vtree: &Vtree, formula: &CnfFormula) -> Result<(), VitriError> {
let indexed = vtree.num_vars() as usize;
if formula.num_vars as usize <= indexed {
return Ok(());
}
for clause in &formula.clauses {
for lit in &clause.literals {
if lit.var.idx() >= indexed {
return Err(VitriError::mismatch(format!(
"vtree indexes {indexed} variables but the formula names DIMACS variable {}; \
the vtree does not belong to this formula",
lit.var.to_dimacs(),
)));
}
}
}
Ok(())
}
pub(crate) const BUILT_FROM_THIS_FORMULA: &str = "vtree was built from this formula";
fn for_each_clause_lca(vtree: &Vtree, formula: &CnfFormula, mut f: impl FnMut(usize, VtreeIdx)) {
for (clause_idx, clause) in formula.clauses.iter().enumerate() {
if let Some(lca) = clause_lca(vtree, clause) {
f(clause_idx, lca);
}
}
}
fn clause_lca(vtree: &Vtree, clause: &Clause) -> Option<VtreeIdx> {
clause
.literals
.iter()
.map(|lit| vtree.leaf_of(lit.var))
.reduce(|a, b| vtree.lca(a, b))
}
fn clause_lca_counts(vtree: &Vtree, formula: &CnfFormula) -> Vec<u32> {
let mut clause_at = vec![0u32; vtree.num_nodes()];
for_each_clause_lca(vtree, formula, |_, lca| clause_at[lca.idx()] += 1);
clause_at
}
#[cfg(test)]
fn clause_lca_buckets(vtree: &Vtree, formula: &CnfFormula) -> (Vec<u32>, Vec<Vec<usize>>) {
(
clause_lca_counts(vtree, formula),
clause_lca_members(vtree, formula),
)
}
fn clause_lca_members(vtree: &Vtree, formula: &CnfFormula) -> Vec<Vec<usize>> {
let mut clauses_at = vec![Vec::new(); vtree.num_nodes()];
for_each_clause_lca(vtree, formula, |clause_idx, lca| {
clauses_at[lca.idx()].push(clause_idx);
});
clauses_at
}
pub(crate) fn clause_lca_nodes(vtree: &Vtree, formula: &CnfFormula) -> (Vec<VtreeIdx>, Vec<u32>) {
let mut per_clause = Vec::with_capacity(formula.clauses.len());
let mut clause_at = vec![0u32; vtree.num_nodes()];
for_each_clause_lca(vtree, formula, |_, lca| {
per_clause.push(lca);
clause_at[lca.idx()] += 1;
});
(per_clause, clause_at)
}
pub(crate) fn vtree_clause_load_per_node(vtree: &Vtree, formula: &CnfFormula) -> Vec<u32> {
clause_lca_counts(vtree, formula)
}
pub fn vtree_cost(vtree: &Vtree, formula: &CnfFormula) -> Result<f64, VitriError> {
Ok(VtreeScores::compute(vtree, formula, None)?.cost)
}
pub fn check_score_env() -> Result<(), VitriError> {
let ranker = agg::model()?;
agg::margin_from_env(ranker.is_some())?;
Ok(())
}
fn log2_sum_exp(values: &[f64]) -> f64 {
let peak = values
.iter()
.copied()
.filter(|&value| value > 0.0)
.reduce(f64::max);
let Some(peak) = peak else {
return 0.0;
};
peak + values
.iter()
.copied()
.filter(|&value| value > 0.0)
.map(|value| 2f64.powf(value - peak))
.sum::<f64>()
.log2()
}
fn cut_width(ctx_in: u32, cross: u32, clause_count: u64) -> f64 {
f64::from(cross) * f64::from(ctx_in) / clause_count.max(1) as f64
}
fn sorted_bounds(ctx_in: u32, ctx_out: u32, cross: u32, is_leaf: bool) -> [u32; 3] {
let mut bounds = [ctx_in, ctx_out, cross];
bounds.sort_unstable();
if is_leaf {
for bound in &mut bounds {
*bound = (*bound).min(1);
}
}
bounds
}
fn separator_terms(
vtree: &Vtree,
ctx_in: &[u32],
ctx_out: &[u32],
cross: &[u32],
clause_count: u64,
) -> (f64, f64, f64, Vec<u32>) {
let mut tight_widths = vec![0u32; vtree.num_nodes()];
let mut second = vec![0u32; vtree.num_nodes()];
let mut cross_width = vec![0u32; vtree.num_nodes()];
let mut width = vec![0f64; vtree.num_nodes()];
for t in vtree.bottomup() {
let i = t.idx();
let is_leaf = vtree.node(t).is_leaf();
let bounds = sorted_bounds(ctx_in[i], ctx_out[i], cross[i], is_leaf);
let leaf_cap = if is_leaf { 1 } else { u32::MAX };
tight_widths[i] = bounds[0];
second[i] = bounds[1];
cross_width[i] = cross[i].min(leaf_cap);
width[i] = if is_leaf {
f64::from(tight_widths[i])
} else {
cut_width(ctx_in[i], cross[i], clause_count)
};
}
let mut tight_terms = Vec::with_capacity(vtree.num_nodes() / 2);
let mut bound_terms = Vec::with_capacity(vtree.num_nodes() / 2);
let mut capped_terms = Vec::with_capacity(vtree.num_nodes() / 2);
let mut pair_cross_terms = Vec::with_capacity(vtree.num_nodes() / 2);
let mut out_terms = Vec::with_capacity(vtree.num_nodes() / 2);
for (t, left, right) in vtree.internal_bottomup() {
let i = t.idx();
tight_terms.push(width[i]);
bound_terms.push(f64::from(tight_widths[i]));
capped_terms
.push(f64::from(tight_widths[i]) + f64::from((second[i] - tight_widths[i]).min(7)));
pair_cross_terms
.push(f64::from(cross_width[left.idx()]) + f64::from(cross_width[right.idx()]));
out_terms.push(f64::from(ctx_out[i]));
}
let tight = log2_sum_exp(&tight_terms);
let tight_bound = log2_sum_exp(&bound_terms);
let capped_gap = (log2_sum_exp(&capped_terms) - tight_bound).max(0.0);
let pair_cross_gap = (log2_sum_exp(&pair_cross_terms) - tight_bound).max(0.0);
let excess = (capped_gap - (1.0 + pair_cross_gap).log2()).max(0.0);
(tight, excess, log2_sum_exp(&out_terms), tight_widths)
}
const UNIQUE_PRESSURE_THRESHOLD: f64 = 7.672_358_059_638_748;
struct ChildBoundaryFeatures {
outside_overlap_top2_mean: f64,
outside_symmetric_difference_max: u32,
tight_unique_sum: Vec<u32>,
}
fn child_boundary_features(
vtree: &Vtree,
tight_widths: &[u32],
outside_widths: &[u32],
sibling_overlap: &[u32],
) -> ChildBoundaryFeatures {
let mut largest_overlap = 0u32;
let mut second_overlap = 0u32;
let mut internal_count = 0u32;
let mut symmetric_difference_max = 0u32;
let mut tight_unique_sum = vec![0u32; vtree.num_nodes()];
for (node, left, right) in vtree.internal_bottomup() {
internal_count += 1;
let overlap = sibling_overlap[node.idx()];
if overlap >= largest_overlap {
second_overlap = largest_overlap;
largest_overlap = overlap;
} else if overlap > second_overlap {
second_overlap = overlap;
}
symmetric_difference_max = symmetric_difference_max
.max(outside_widths[left.idx()] + outside_widths[right.idx()] - 2 * overlap);
let tight_overlap = overlap
.min(tight_widths[left.idx()])
.min(tight_widths[right.idx()]);
tight_unique_sum[node.idx()] =
tight_widths[left.idx()] + tight_widths[right.idx()] - tight_overlap;
}
let outside_overlap_top2_mean = match internal_count {
0 => 0.0,
1 => f64::from(largest_overlap),
_ => f64::from(largest_overlap + second_overlap) / 2.0,
};
ChildBoundaryFeatures {
outside_overlap_top2_mean,
outside_symmetric_difference_max: symmetric_difference_max,
tight_unique_sum,
}
}
fn successor_guard_correction(
tight_unique_pressure_max: f64,
outside_overlap_top2_mean: f64,
outside_symmetric_difference_max: u32,
) -> f64 {
0.55 * (tight_unique_pressure_max - UNIQUE_PRESSURE_THRESHOLD).clamp(0.0, 0.25)
+ 1.5 * (37.0 - outside_overlap_top2_mean).clamp(0.0, 1.0)
+ 3.84 * (outside_overlap_top2_mean - 22.5).clamp(0.0, 1.0)
+ 1.5 * (63.0 - f64::from(outside_symmetric_difference_max)).clamp(0.0, 1.0)
}
fn node_depths(vtree: &Vtree) -> Vec<u32> {
let mut depth = vec![0u32; vtree.num_nodes()];
for node in vtree.bottomup().rev() {
if !vtree.node(node).is_leaf() {
let (left, right) = vtree.children(node);
depth[left.idx()] = depth[node.idx()] + 1;
depth[right.idx()] = depth[node.idx()] + 1;
}
}
depth
}
fn vtree_depth(vtree: &Vtree) -> u32 {
node_depths(vtree).into_iter().max().unwrap_or(0)
}
fn context_direction_sums(vtree: &Vtree, ctx_in: &[u32]) -> (f64, f64) {
let mut left_terms = Vec::with_capacity(vtree.num_nodes() / 2);
let mut right_terms = Vec::with_capacity(vtree.num_nodes() / 2);
for (_, left, right) in vtree.internal_bottomup() {
left_terms.push(f64::from(ctx_in[left.idx()]));
right_terms.push(f64::from(ctx_in[right.idx()]));
}
(log2_sum_exp(&left_terms), log2_sum_exp(&right_terms))
}
fn directional_context_excess(vtree: &Vtree, ctx_in: &[u32], depth: u32) -> f64 {
if 5 * u64::from(depth) > u64::from(vtree.num_leaves()) + 1 {
return 0.0;
}
let (left, right) = context_direction_sums(vtree, ctx_in);
(left - right - 3.0).max(0.0)
}
fn output_gap_bits(tight: f64, outside: f64) -> f64 {
(1.0 + (outside - tight - 12.0).max(0.0)).log2()
}
fn extreme_chain_guard(leaves: u32, depth: u32) -> f64 {
let leaves = f64::from(leaves.max(2));
let depth_ratio = (f64::from(depth) / (leaves - 1.0)).min(1.0);
let remaining = (1.0 / leaves).max(1.0 - depth_ratio);
(-remaining.log2() - 2.0).max(0.0)
}
fn extreme_local_join_guard(join_excess: f64) -> f64 {
(join_excess - 12.0).max(0.0)
}
struct SubtreeTables {
clauses: Vec<u64>,
leaves: Vec<u32>,
height: Vec<u32>,
}
fn subtree_tables(vtree: &Vtree, clause_at: &[u32]) -> SubtreeTables {
let mut clauses = vec![0u64; vtree.num_nodes()];
let mut leaves = vec![0u32; vtree.num_nodes()];
let mut height = vec![0u32; vtree.num_nodes()];
for t in vtree.bottomup() {
let i = t.idx();
if vtree.node(t).is_leaf() {
clauses[i] = u64::from(clause_at[i]);
leaves[i] = 1;
continue;
}
let (left, right) = vtree.children(t);
clauses[i] = u64::from(clause_at[i]) + clauses[left.idx()] + clauses[right.idx()];
leaves[i] = leaves[left.idx()] + leaves[right.idx()];
height[i] = 1 + height[left.idx()].max(height[right.idx()]);
}
SubtreeTables {
clauses,
leaves,
height,
}
}
fn clause_load_cost(vtree: &Vtree, clause_at: &[u32]) -> f64 {
let subtree = subtree_tables(vtree, clause_at);
let mut child_products = 0.0;
let mut scope = 0.0;
for (t, left, right) in vtree.internal_bottomup() {
child_products += subtree.clauses[left.idx()] as f64 * subtree.clauses[right.idx()] as f64;
scope += f64::from(clause_at[t.idx()]) * f64::from(subtree.leaves[t.idx()].ilog2());
}
let max_load = f64::from(max_from_counts(clause_at));
max_load.powi(3) + child_products + scope
}
fn maximum_matching_size(adjacency: &[Vec<usize>]) -> u32 {
let mut pair_left = vec![None; adjacency.len()];
let mut pair_right = HashMap::new();
let mut left_seen = vec![0u32; adjacency.len()];
let mut right_seen = HashMap::new();
let mut parent_right = HashMap::new();
let mut visit = 0u32;
let mut size = 0u32;
for start in 0..adjacency.len() {
if pair_left[start].is_some() {
continue;
}
visit = visit.checked_add(1).unwrap_or_else(|| {
left_seen.fill(0);
right_seen.clear();
1
});
let mut queue = VecDeque::from([start]);
left_seen[start] = visit;
let mut endpoint = None;
'search: while let Some(left) = queue.pop_front() {
for &right in &adjacency[left] {
if right_seen.get(&right) == Some(&visit) {
continue;
}
right_seen.insert(right, visit);
parent_right.insert(right, left);
match pair_right.get(&right).copied() {
None => {
endpoint = Some(right);
break 'search;
}
Some(mate) if left_seen[mate] != visit => {
left_seen[mate] = visit;
queue.push_back(mate);
}
Some(_) => {}
}
}
}
let Some(mut right) = endpoint else {
continue;
};
loop {
let left = parent_right[&right];
let previous = pair_left[left];
pair_left[left] = Some(right);
pair_right.insert(right, left);
let Some(previous) = previous else {
break;
};
right = previous;
}
size += 1;
}
size
}
fn subtree_intervals(vtree: &Vtree) -> (Vec<u32>, Vec<u32>) {
let mut entry = vec![0u32; vtree.num_nodes()];
let mut exit = vec![0u32; vtree.num_nodes()];
let mut next = 0u32;
let mut stack = vec![(vtree.root(), false)];
while let Some((node, leaving)) = stack.pop() {
if leaving {
exit[node.idx()] = next;
continue;
}
entry[node.idx()] = next;
next += 1;
stack.push((node, true));
if !vtree.node(node).is_leaf() {
let (left, right) = vtree.children(node);
stack.push((right, false));
stack.push((left, false));
}
}
(entry, exit)
}
#[cfg(test)]
fn local_join_match_excess(
vtree: &Vtree,
formula: &CnfFormula,
clauses_at: &[Vec<usize>],
clause_count: u64,
) -> f64 {
local_join_features(
vtree,
formula,
clauses_at,
clause_count,
true,
&vec![0; vtree.num_nodes()],
)
.0
}
fn local_join_features(
vtree: &Vtree,
formula: &CnfFormula,
clauses_at: &[Vec<usize>],
clause_count: u64,
shallow: bool,
tight_unique_sum: &[u32],
) -> (f64, f64) {
if clauses_at.is_empty() {
return (0.0, 0.0);
}
let (entry, exit) = subtree_intervals(vtree);
let mut peak_excess = 0.0f64;
let mut tight_unique_pressure_max = 0.0f64;
for (t, left, _) in vtree.internal_bottomup() {
let clause_ids = &clauses_at[t.idx()];
if clause_ids.is_empty() {
continue;
}
let load = clause_ids.len() as u64;
let unique_scale = (1.0 + f64::from(tight_unique_sum[t.idx()])).log2();
let density_upper = load as f64 * load as f64 / clause_count.max(1) as f64;
let can_clear_join = shallow && density_upper > 4.0;
let can_clear_pressure = density_upper * unique_scale > UNIQUE_PRESSURE_THRESHOLD;
if !can_clear_join && !can_clear_pressure {
continue;
}
let mut left_adjacency = Vec::with_capacity(clause_ids.len());
let mut right_adjacency = Vec::with_capacity(clause_ids.len());
for &clause_idx in clause_ids {
let mut left_vars = Vec::new();
let mut right_vars = Vec::new();
for lit in &formula.clauses[clause_idx].literals {
let var = lit.var.idx();
let leaf = vtree.leaf_of(lit.var);
if entry[left.idx()] <= entry[leaf.idx()] && entry[leaf.idx()] < exit[left.idx()] {
left_vars.push(var);
} else {
right_vars.push(var);
}
}
left_vars.sort_unstable();
left_vars.dedup();
right_vars.sort_unstable();
right_vars.dedup();
left_adjacency.push(left_vars);
right_adjacency.push(right_vars);
}
let matching =
maximum_matching_size(&left_adjacency).min(maximum_matching_size(&right_adjacency));
let density = f64::from(matching) * clause_ids.len() as f64 / clause_count.max(1) as f64;
if shallow {
peak_excess = peak_excess.max(density - 4.0);
}
tight_unique_pressure_max = tight_unique_pressure_max.max(density * unique_scale);
}
(peak_excess, tight_unique_pressure_max)
}
struct UnifiedCostTables<'a> {
clause_at: &'a [u32],
ctx_in: &'a [u32],
ctx_out: &'a [u32],
sibling_overlap: &'a [u32],
cross: &'a [u32],
}
pub const COST_TERM_NAMES: [&str; 11] = [
"tight",
"excess_half",
"clause_load_bits",
"high_load_25",
"chain_3_40",
"join_neg_half",
"directional_half",
"output_gap_16",
"extreme_chain_4",
"extreme_join_32",
"successor_guard",
];
pub fn vtree_cost_terms(vtree: &Vtree, formula: &CnfFormula) -> Result<[f64; 11], VitriError> {
covered_by(vtree, formula)?;
let tables = tables::Tables::build(vtree, formula, false, false);
Ok(unified_cost_terms(
vtree,
formula,
tables.cost_tables(),
stddev_from_counts(tables.clause_at()),
vtree_depth(vtree),
))
}
pub(in crate::score) fn unified_cost_terms(
vtree: &Vtree,
formula: &CnfFormula,
tables: UnifiedCostTables<'_>,
load_stddev: f64,
depth: u32,
) -> [f64; 11] {
let clause_count: u64 = tables.clause_at.iter().map(|&load| u64::from(load)).sum();
let (tight, excess, outside, tight_widths) = separator_terms(
vtree,
tables.ctx_in,
tables.ctx_out,
tables.cross,
clause_count,
);
if tight == 0.0 {
return [0.0; 11];
}
let child_boundaries =
child_boundary_features(vtree, &tight_widths, tables.ctx_out, tables.sibling_overlap);
let clause_load_cost = clause_load_cost(vtree, tables.clause_at);
let leaves = f64::from(vtree.num_leaves());
let chain = (1.0 + (5.0 * f64::from(depth) - leaves - 1.0).max(0.0)).log2();
let high_load = (1.0 + load_stddev).log2() * (tight - 16.0).max(0.0);
let shallow = 5 * u64::from(depth) <= u64::from(vtree.num_leaves()) + 1;
let needs_matching = vtree.internal_bottomup().any(|(node, _, _)| {
let load = f64::from(tables.clause_at[node.idx()]);
let density_upper = load * load / clause_count.max(1) as f64;
let unique_scale = (1.0 + f64::from(child_boundaries.tight_unique_sum[node.idx()])).log2();
(shallow && density_upper > 4.0) || density_upper * unique_scale > UNIQUE_PRESSURE_THRESHOLD
});
let clauses_at = if needs_matching {
clause_lca_members(vtree, formula)
} else {
Vec::new()
};
let (join, tight_unique_pressure) = local_join_features(
vtree,
formula,
&clauses_at,
clause_count,
shallow,
&child_boundaries.tight_unique_sum,
);
let directional_context = directional_context_excess(vtree, tables.ctx_in, depth);
let output_gap = output_gap_bits(tight, outside);
let extreme_chain = extreme_chain_guard(vtree.num_leaves(), depth);
let extreme_join = extreme_local_join_guard(join);
let successor_guard = successor_guard_correction(
tight_unique_pressure,
child_boundaries.outside_overlap_top2_mean,
child_boundaries.outside_symmetric_difference_max,
);
[
tight,
excess / 2.0,
9.0 * (1.0 + clause_load_cost).log2() / 5.0,
high_load / 25.0,
3.0 * chain / 40.0,
-join / 2.0,
directional_context / 2.0,
8.0 * output_gap / 5.0,
4.0 * extreme_chain,
32.0 * extreme_join,
successor_guard,
]
}
pub(crate) fn vtree_max_clause_load(vtree: &Vtree, formula: &CnfFormula) -> u32 {
max_from_counts(&clause_lca_counts(vtree, formula))
}
fn max_from_counts(clause_at: &[u32]) -> u32 {
clause_at.iter().copied().max().unwrap_or(0)
}
fn stddev_from_counts(clause_at: &[u32]) -> f64 {
load_stats(clause_at, |_| true).stddev
}
pub(crate) struct LoadStats {
pub(crate) mean: f64,
pub(crate) stddev: f64,
pub(crate) count: usize,
}
pub(crate) fn load_stats(loads: &[u32], keep: impl Fn(VtreeIdx) -> bool) -> LoadStats {
let mut sum: f64 = 0.0;
let mut sum_sq: f64 = 0.0;
let mut count: usize = 0;
for (idx, &load) in loads.iter().enumerate() {
if load > 0 && keep(VtreeIdx(idx as u32)) {
sum += load as f64;
sum_sq += (load as f64) * (load as f64);
count += 1;
}
}
if count == 0 {
return LoadStats {
mean: 0.0,
stddev: 0.0,
count: 0,
};
}
let mean = sum / count as f64;
let stddev = if count == 1 {
0.0
} else {
((sum_sq - sum * mean) / (count - 1) as f64).max(0.0).sqrt()
};
LoadStats {
mean,
stddev,
count,
}
}
pub(crate) fn vtree_context_width_per_node(
vtree: &Vtree,
formula: &CnfFormula,
show: Option<&crate::cnf::ShowMask>,
) -> Vec<u32> {
let high_lca = clause_high_lca(vtree, formula);
context_width_from_high_lca(vtree, &high_lca, show)
}
fn context_width_from_high_lca(
vtree: &Vtree,
high_lca: &[Option<VtreeIdx>],
show: Option<&crate::cnf::ShowMask>,
) -> Vec<u32> {
let mut ctx = vec![0u32; vtree.num_nodes()];
for (vi, &lca) in high_lca.iter().enumerate() {
if let Some(mask) = show
&& !mask.as_slice().get(vi).copied().unwrap_or(false)
{
continue;
}
if let Some(l) = lca {
let leaf = vtree.leaf_of(VarId(vi as u32));
let mut cur = vtree.node(leaf).parent();
while let Some(node) = cur {
if node == l {
break; }
ctx[node.idx()] += 1;
cur = vtree.node(node).parent();
}
}
}
ctx
}
#[cfg(test)]
pub(crate) fn vtree_outside_context_width_per_node(
vtree: &Vtree,
formula: &CnfFormula,
) -> Vec<u32> {
outside_context_tables(vtree, formula).widths
}
struct OutsideContextTables {
widths: Vec<u32>,
sibling_overlap: Vec<u32>,
}
fn outside_context_tables(vtree: &Vtree, formula: &CnfFormula) -> OutsideContextTables {
let n_vars = vtree.num_vars() as usize;
let (pos, neg) = crate::cnf::occ::occurrence_lists(&formula.clauses, n_vars);
let nn = vtree.num_nodes();
let mut ctx_out = vec![0u32; nn];
let mut sibling_overlap = vec![0u32; nn];
let mut stamp: Vec<u32> = vec![u32::MAX; nn];
let mut outside_stamp: Vec<u32> = vec![u32::MAX; nn];
for (v, (in_pos, in_neg)) in pos.iter().zip(&neg).enumerate() {
if in_pos.is_empty() && in_neg.is_empty() {
continue;
}
let v_id = v as u32;
let mut cur = Some(vtree.leaf_of(VarId(v_id)));
while let Some(node) = cur {
stamp[node.idx()] = v_id;
cur = vtree.node(node).parent();
}
for &ci in in_pos.iter().chain(in_neg) {
for lit in &formula.clauses[ci].literals {
if lit.var.idx() == v {
continue;
}
let mut cur = Some(vtree.leaf_of(lit.var));
while let Some(node) = cur {
if stamp[node.idx()] == v_id {
break;
}
stamp[node.idx()] = v_id;
ctx_out[node.idx()] += 1;
outside_stamp[node.idx()] = v_id;
if let Some(parent) = vtree.node(node).parent() {
let (left, right) = vtree.children(parent);
let sibling = if node == left { right } else { left };
if outside_stamp[sibling.idx()] == v_id {
sibling_overlap[parent.idx()] += 1;
}
}
cur = vtree.node(node).parent();
}
}
}
}
OutsideContextTables {
widths: ctx_out,
sibling_overlap,
}
}
pub(crate) fn vtree_crossing_clauses_per_node(vtree: &Vtree, formula: &CnfFormula) -> Vec<u32> {
let nn = vtree.num_nodes();
let mut cross = vec![0u32; nn];
let mut stamp: Vec<usize> = vec![usize::MAX; nn];
for (ci, clause) in formula.clauses.iter().enumerate() {
if clause.literals.len() < 2 {
continue;
}
let lca = clause_lca(vtree, clause).expect("a clause with two literals has an LCA");
stamp[lca.idx()] = ci;
for lit in &clause.literals {
let mut cur = vtree.leaf_of(lit.var);
while stamp[cur.idx()] != ci {
stamp[cur.idx()] = ci;
cross[cur.idx()] += 1;
cur = vtree
.node(cur)
.parent()
.expect("a node below the clause LCA has a parent");
}
}
}
cross
}
fn clause_high_lca(vtree: &Vtree, formula: &CnfFormula) -> Vec<Option<VtreeIdx>> {
let n_vars = vtree.num_vars() as usize;
let mut high_lca: Vec<Option<VtreeIdx>> = vec![None; n_vars];
for clause in &formula.clauses {
if clause.literals.len() < 2 {
continue; }
let lca = clause_lca(vtree, clause).expect("a clause with two literals has an LCA");
let lpos = vtree.topo_pos(lca);
for lit in &clause.literals {
let vi = lit.var.idx();
let replace = match high_lca[vi] {
Some(cur) => lpos > vtree.topo_pos(cur),
None => true,
};
if replace {
high_lca[vi] = Some(lca);
}
}
}
high_lca
}
#[derive(Clone, Copy, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct VtreeScores {
pub clause_load_stddev: f64,
pub max_clause_load: u32,
pub peak_context_width_all: u32,
pub peak_context_width_show: Option<u32>,
pub cost: f64,
}
impl VtreeScores {
pub fn compute(
vtree: &Vtree,
formula: &CnfFormula,
show_mask: Option<&crate::cnf::ShowMask>,
) -> Result<Self, VitriError> {
covered_by(vtree, formula)?;
let clause_at = clause_lca_counts(vtree, formula);
let high_lca = clause_high_lca(vtree, formula);
let ctx_in = context_width_from_high_lca(vtree, &high_lca, None);
let outside = outside_context_tables(vtree, formula);
let cross = vtree_crossing_clauses_per_node(vtree, formula);
let peak_show = show_mask.map(|m| {
context_width_from_high_lca(vtree, &high_lca, Some(m))
.into_iter()
.max()
.unwrap_or(0)
});
Ok(Self::from_tables(
vtree,
formula,
UnifiedCostTables {
clause_at: &clause_at,
ctx_in: &ctx_in,
ctx_out: &outside.widths,
sibling_overlap: &outside.sibling_overlap,
cross: &cross,
},
peak_show,
)
.0)
}
fn from_tables(
vtree: &Vtree,
formula: &CnfFormula,
tables: UnifiedCostTables<'_>,
peak_context_width_show: Option<u32>,
) -> (Self, [f64; 11]) {
let clause_load_stddev = stddev_from_counts(tables.clause_at);
let max_clause_load = max_from_counts(tables.clause_at);
let peak_context_width_all = tables.ctx_in.iter().copied().max().unwrap_or(0);
let terms = unified_cost_terms(
vtree,
formula,
tables,
clause_load_stddev,
vtree_depth(vtree),
);
let cost = terms.iter().sum();
(
Self {
clause_load_stddev,
max_clause_load,
peak_context_width_all,
peak_context_width_show,
cost,
},
terms,
)
}
}
#[cfg(test)]
mod tests;
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct StructureProfile {
pub clause_width_cv: f64,
pub var_occurrence_cv: f64,
pub coloring_like: bool,
}
impl StructureProfile {
pub fn from_coefficients(clause_width_cv: f64, var_occurrence_cv: f64) -> Self {
StructureProfile {
clause_width_cv,
var_occurrence_cv,
coloring_like: crate::cnf::stats::coloring_like_predicate(
var_occurrence_cv,
clause_width_cv,
),
}
}
pub fn measure(formula: &CnfFormula) -> Self {
let clause_width_cv = crate::cnf::stats::clause_width_cv(formula);
let var_occurrence_cv = crate::cnf::stats::var_occurrence_cv(formula);
StructureProfile::from_coefficients(clause_width_cv, var_occurrence_cv)
}
}