use std::cell::OnceCell;
use std::collections::HashSet;
use crate::cnf::CnfFormula;
use crate::vtree::{VarId, Vtree, VtreeArena, VtreeIdx};
use super::super::TreeDecomposition;
use super::super::td_parse::primal_adjacency;
use super::combiners::{combine_edge_aligned, combine_hypergraph_bisect, combine_into_balanced};
use super::meta::BagMetadata;
use super::reading::{Binarization, FixedReading, Place, RootPick};
#[derive(Clone, Copy)]
pub(crate) struct ConversionInput<'a> {
pub td: &'a TreeDecomposition,
pub num_vars: u32,
pub formula: Option<&'a CnfFormula>,
pub effort_scale: f64,
}
pub(super) struct Converter<'a> {
pub(super) input: ConversionInput<'a>,
primal_adj: OnceCell<Vec<Vec<u32>>>,
}
impl<'a> Converter<'a> {
pub(super) fn new(input: ConversionInput<'a>) -> Self {
Self {
input,
primal_adj: OnceCell::new(),
}
}
pub(super) fn build(&self, reading: FixedReading) -> (Vtree, BagMetadata) {
let ConversionInput {
td,
num_vars,
formula,
effort_scale,
} = self.input;
let primal_adj = formula
.filter(|_| reading.place == Place::Deep || reading.binarize == Binarization::Edge)
.map(|f| {
self.primal_adj
.get_or_init(|| primal_adjacency(f, num_vars))
.as_slice()
});
let n = td.bags().len();
let chosen = root_bags(td, reading.root);
let forest = td
.rooted_forest(chosen.iter().copied())
.expect("conversion roots are bag indices");
let order = forest.order();
let parent_td = forest.parents();
let depth = forest.depths();
let component_roots = forest.component_roots();
let var_bag = assign_var_bags(td, num_vars, order, depth, reading.place, primal_adj);
let meta = BagMetadata::from_assignment(num_vars, &var_bag, order, n, td.treewidth());
let mut vars_at: Vec<Vec<u32>> = vec![Vec::new(); n];
let mut in_any_bag = vec![false; num_vars as usize];
for v in 0..num_vars as usize {
if var_bag[v] != usize::MAX {
vars_at[var_bag[v]].push(v as u32);
in_any_bag[v] = true;
}
}
let mut nodes = VtreeArena::new();
let mut td_vtree_idx: Vec<Option<VtreeIdx>> = vec![None; n];
let track_assigned_vars = reading.binarize == Binarization::Hypergraph && formula.is_some();
let mut td_vars: Vec<Vec<u32>> = if track_assigned_vars {
vec![Vec::new(); n]
} else {
Vec::new()
};
let track_bag_vars = reading.binarize == Binarization::Edge && formula.is_some();
let mut td_bag_vars: Vec<HashSet<u32>> = if track_bag_vars {
vec![HashSet::new(); n]
} else {
Vec::new()
};
for &t in order.iter().rev() {
let mut child_items: Vec<VtreeIdx> = Vec::new();
let mut child_var_sets: Vec<Vec<u32>> = Vec::new();
let mut child_bag_var_sets: Vec<HashSet<u32>> = Vec::new();
for &nb in &td.adjacency()[t] {
if Some(nb) != parent_td[t]
&& let Some(child_idx) = td_vtree_idx[nb]
{
child_items.push(child_idx);
if track_assigned_vars {
child_var_sets.push(std::mem::take(&mut td_vars[nb]));
}
if track_bag_vars {
child_bag_var_sets.push(std::mem::take(&mut td_bag_vars[nb]));
}
}
}
let mut var_items: Vec<VtreeIdx> = Vec::new();
for &v in &vars_at[t] {
let idx = nodes.leaf(VarId(v));
var_items.push(idx);
}
let mut items = child_items.clone();
items.extend_from_slice(&var_items);
if track_assigned_vars {
let mut all_vars: Vec<u32> = Vec::new();
for cv in &child_var_sets {
all_vars.extend_from_slice(cv);
}
all_vars.extend_from_slice(&vars_at[t]);
td_vars[t] = all_vars;
}
td_vtree_idx[t] = if items.is_empty() {
None
} else {
let bag = BagItems {
items: &items,
child_items: &child_items,
child_var_sets: &child_var_sets,
child_bag_var_sets: &child_bag_var_sets,
var_items: &var_items,
vars_here: &vars_at[t],
};
Some(combine_bag(
&bag,
reading.binarize,
formula,
effort_scale,
primal_adj.unwrap_or_default(),
&mut nodes,
))
};
if track_bag_vars {
let mut bag_union = child_bag_var_sets
.iter()
.enumerate()
.max_by_key(|(_, vars)| vars.len())
.map(|(index, _)| index)
.map(|index| child_bag_var_sets.swap_remove(index))
.unwrap_or_default();
for child_vars in child_bag_var_sets {
bag_union.extend(child_vars);
}
bag_union.extend(
td.bags()[t]
.vertices()
.iter()
.copied()
.filter(|&v| v < num_vars),
);
td_bag_vars[t] = bag_union;
}
}
let mut top_items: Vec<VtreeIdx> = Vec::new();
for &cr in component_roots {
if let Some(root_idx) = td_vtree_idx[cr] {
top_items.push(root_idx);
}
}
for (v, &bagged) in in_any_bag.iter().enumerate() {
if !bagged {
let idx = nodes.leaf(VarId(v as u32));
top_items.push(idx);
}
}
assert!(!top_items.is_empty(), "td_to_vtree: no variables found");
let root = combine_into_balanced(&top_items, &mut nodes);
(Vtree::from_nodes(nodes.into_nodes(), root, num_vars), meta)
}
}
pub(super) fn root_bags(td: &TreeDecomposition, root: RootPick) -> Vec<usize> {
let first = || {
td.rooted_forest(0..td.bags().len())
.expect("all generated roots are bag indices")
.component_roots()
.to_vec()
};
match root {
RootPick::Leaf(bag) => vec![bag],
RootPick::First => first(),
RootPick::Centroid => first().iter().map(|&cr| find_centroid(td, cr)).collect(),
}
}
fn assign_var_bags(
td: &TreeDecomposition,
num_vars: u32,
order: &[usize],
depth: &[usize],
place: Place,
primal_adj: Option<&[Vec<u32>]>,
) -> Vec<usize> {
let mut var_bag = vec![usize::MAX; num_vars as usize];
match place {
Place::Deep => {
let mut var_max_depth = vec![0usize; num_vars as usize];
for &bag_idx in order {
for &v in td.bags()[bag_idx].vertices() {
if (v as usize) < num_vars as usize {
var_bag[v as usize] = bag_idx;
var_max_depth[v as usize] = depth[bag_idx];
}
}
}
if let Some(primal_adj) = primal_adj {
apply_cooc_tiebreak(
&BagWalk { td, order, depth },
primal_adj,
num_vars,
&mut var_bag,
&var_max_depth,
);
}
}
Place::Shallow => {
for (bag_idx, bag) in td.bags().iter().enumerate() {
for &v in bag.vertices() {
if (v as usize) < num_vars as usize {
let cur = var_bag[v as usize];
if cur == usize::MAX || depth[bag_idx] < depth[cur] {
var_bag[v as usize] = bag_idx;
}
}
}
}
}
}
var_bag
}
struct BagItems<'a> {
items: &'a [VtreeIdx],
child_items: &'a [VtreeIdx],
child_var_sets: &'a [Vec<u32>],
child_bag_var_sets: &'a [HashSet<u32>],
var_items: &'a [VtreeIdx],
vars_here: &'a [u32],
}
fn combine_bag(
bag: &BagItems<'_>,
binarize: Binarization,
formula: Option<&CnfFormula>,
effort_scale: f64,
edge_primal_adj: &[Vec<u32>],
nodes: &mut VtreeArena,
) -> VtreeIdx {
let items = bag.items;
match (binarize, formula) {
(Binarization::Hypergraph, Some(formula)) => {
let mut item_vars: Vec<Vec<u32>> = bag.child_var_sets.to_vec();
for &v in bag.vars_here {
item_vars.push(vec![v]);
}
debug_assert_eq!(
item_vars.len(),
items.len(),
"item_vars len {} != items len {} (children={}, vars={})",
item_vars.len(),
items.len(),
bag.child_var_sets.len(),
bag.vars_here.len()
);
combine_hypergraph_bisect(items, &item_vars, formula, effort_scale, nodes)
}
(Binarization::Edge, Some(_)) => combine_edge_aligned(
bag.child_items,
bag.child_bag_var_sets,
bag.var_items,
bag.vars_here,
edge_primal_adj,
nodes,
),
_ => combine_into_balanced(items, nodes),
}
}
pub(super) fn find_centroid(td: &TreeDecomposition, start: usize) -> usize {
let adj = td.adjacency();
debug_assert!(
!adj.is_empty(),
"find_centroid requires a non-empty decomposition"
);
let forest = td
.rooted_forest([start])
.expect("centroid start is a bag index");
let all_order = forest.order();
let parent = forest.parents();
let component_size = all_order
.iter()
.skip(1)
.position(|&bag| parent[bag].is_none())
.map_or(all_order.len(), |offset| offset + 1);
let order = &all_order[..component_size];
if component_size <= 2 {
return start;
}
let mut subtree_size = vec![1usize; adj.len()];
for &t in order.iter().rev() {
if let Some(parent) = parent[t] {
let child_size = subtree_size[t];
subtree_size[parent] += child_size;
}
}
let mut best_node = start;
let mut best_max = component_size;
for &t in order {
let mut max_part = component_size - subtree_size[t]; for &nb in &adj[t] {
if Some(nb) != parent[t] {
max_part = max_part.max(subtree_size[nb]);
}
}
if max_part < best_max {
best_max = max_part;
best_node = t;
}
}
best_node
}
struct BagWalk<'a> {
td: &'a TreeDecomposition,
order: &'a [usize],
depth: &'a [usize],
}
fn apply_cooc_tiebreak(
walk: &BagWalk<'_>,
primal_adj: &[Vec<u32>],
num_vars: u32,
var_bag: &mut [usize],
var_max_depth: &[usize],
) {
let nv = num_vars as usize;
let mut best_score_u: Vec<u32> = vec![0u32; nv];
let mut in_bag = vec![false; nv];
for &bag_idx in walk.order {
let d = walk.depth[bag_idx];
let bag_vars: Vec<u32> = walk.td.bags()[bag_idx]
.vertices()
.iter()
.copied()
.filter(|&v| (v as usize) < nv)
.collect();
if bag_vars.is_empty() {
continue;
}
for &v in &bag_vars {
in_bag[v as usize] = true;
}
for &v in &bag_vars {
let vi = v as usize;
if d == var_max_depth[vi] {
let score: u32 = primal_adj[vi]
.iter()
.filter(|&&u| in_bag[u as usize])
.count() as u32;
if score >= best_score_u[vi] {
best_score_u[vi] = score;
var_bag[vi] = bag_idx;
}
}
}
for &v in &bag_vars {
in_bag[v as usize] = false;
}
}
}