use std::time::Instant;
use crate::budget::expired;
use crate::cnf::CnfFormula;
use crate::vtree::{VarId, Vtree, VtreeArena, VtreeIdx};
#[derive(Clone, Copy)]
pub(crate) struct BisectDials {
pub(crate) imbalance: f64,
pub(crate) base_seed: u64,
pub(crate) effort_scale: f64,
}
pub(crate) const SOLVER_FALLBACK_VARS: usize = 32;
pub(crate) fn minfill_subtree(
vars: &[u32],
formula: &CnfFormula,
nodes: &mut VtreeArena,
) -> VtreeIdx {
let edges = super::td_parse::primal_edges_on_subset(formula, vars);
minfill_subtree_from_local_edges(vars, vars.len() as u32, &edges, nodes)
}
fn minfill_subtree_from_local_edges(
vars: &[u32],
n: u32,
local_edges: &[(u32, u32)],
nodes: &mut VtreeArena,
) -> VtreeIdx {
let td = super::goatd::minfill_td_from_edges(n, local_edges, super::INTERNAL_ELIMINATION_SEED);
let sub_vtree = super::td_to_vtree::td_to_vtree(&td, n);
nodes.graft(&sub_vtree, |local| VarId(vars[local.0 as usize]))
}
pub(crate) struct Bisection {
pub(crate) left: Vec<u32>,
pub(crate) right: Vec<u32>,
}
impl Bisection {
pub(crate) fn from_side_bits(vars: &[u32], bits: &[u8]) -> Option<Bisection> {
if bits.len() != vars.len() {
return None;
}
let mut left = Vec::new();
let mut right = Vec::new();
for (&var, &side) in vars.iter().zip(bits) {
if side == 0 {
left.push(var);
} else {
right.push(var);
}
}
(!left.is_empty() && !right.is_empty()).then_some(Bisection { left, right })
}
}
pub(crate) trait BisectionSolver {
fn partition(
&mut self,
vars: &[u32],
formula: &CnfFormula,
) -> Result<Option<Bisection>, String>;
fn deadline(&self) -> Option<Instant> {
None
}
fn minfill_cutoff(&self) -> usize {
SOLVER_FALLBACK_VARS
}
fn refine_subtree(
&mut self,
vars: &[u32],
formula: &CnfFormula,
nodes: &mut VtreeArena,
checkpoint: usize,
root: VtreeIdx,
) -> Option<VtreeIdx> {
let _ = (vars, formula, nodes, checkpoint, root);
None
}
}
pub(crate) fn run_bisection<S: BisectionSolver>(
formula: &CnfFormula,
solver: &mut S,
) -> Result<std::sync::Arc<Vtree>, String> {
let num_vars = formula.num_vars;
let all_vars: Vec<u32> = (0..num_vars).collect();
let mut nodes = VtreeArena::new();
let root = bisect_recursive_generic(&all_vars, formula, &mut nodes, solver)?;
Ok(std::sync::Arc::new(Vtree::from_nodes(
nodes.into_nodes(),
root,
num_vars,
)))
}
pub(crate) fn bisect_recursive_generic<S: BisectionSolver>(
vars: &[u32],
formula: &CnfFormula,
nodes: &mut VtreeArena,
solver: &mut S,
) -> Result<VtreeIdx, String> {
if expired(solver.deadline()) {
return Err("bisection vtree construction timed out".to_string());
}
if vars.len() == 1 {
let idx = nodes.leaf(VarId(vars[0]));
return Ok(idx);
}
if vars.len() == 2 {
let l_idx = nodes.leaf(VarId(vars[0]));
let r_idx = nodes.leaf(VarId(vars[1]));
let idx = nodes.internal(l_idx, r_idx);
return Ok(idx);
}
let cutoff = solver.minfill_cutoff();
if cutoff > 0 && vars.len() <= cutoff {
return Ok(minfill_subtree(vars, formula, nodes));
}
let Some(split) = solver.partition(vars, formula)? else {
if vars.len() <= 256 {
return Ok(minfill_subtree(vars, formula, nodes));
}
let mid = vars.len() / 2;
let l = bisect_recursive_generic(&vars[..mid], formula, nodes, solver)?;
let r = bisect_recursive_generic(&vars[mid..], formula, nodes, solver)?;
let idx = nodes.internal(l, r);
return Ok(idx);
};
let checkpoint = nodes.len();
let l = bisect_recursive_generic(&split.left, formula, nodes, solver)?;
let r = bisect_recursive_generic(&split.right, formula, nodes, solver)?;
let idx = nodes.internal(l, r);
Ok(solver
.refine_subtree(vars, formula, nodes, checkpoint, idx)
.unwrap_or(idx))
}