use crate::graph::Graph;
#[derive(Clone, Debug, PartialEq)]
pub struct Params {
pub max_nodes: u64,
pub incumbent: Option<Vec<i8>>,
}
impl Default for Params {
fn default() -> Self {
Params { max_nodes: 20_000_000, incumbent: None }
}
}
#[derive(Clone, Debug)]
pub struct Outcome {
pub state: Vec<i8>,
pub energy: f64,
pub proved_optimal: bool,
pub nodes: u64,
pub pruned: u64,
pub hit_limit: bool,
pub slack: f64,
}
pub fn solve(g: &Graph, p: &Params) -> Outcome {
let n = g.n;
if n == 0 {
return Outcome {
state: Vec::new(),
energy: 0.0,
proved_optimal: true,
nodes: 0,
pruned: 0,
hit_limit: false,
slack: 0.0,
};
}
let mut order: Vec<usize> = (0..n).collect();
order.sort_by_key(|&i| core::cmp::Reverse(g.offset[i + 1] - g.offset[i]));
let total: f64 = g.w.iter().map(|v| v.abs()).sum::<f64>() + g.h.iter().map(|v| v.abs()).sum::<f64>();
let slack = (total * n as f64 * f64::EPSILON * 8.0).max(1e-12);
let mut best_state: Vec<i8> = p
.incumbent
.as_ref()
.filter(|s| s.len() == n && s.iter().all(|&v| v == 1 || v == -1))
.cloned()
.unwrap_or_else(|| vec![1i8; n]);
let mut best = g.energy(&best_state);
let gauge_fixed = g.h.iter().all(|&h| h == 0.0);
let mut st = Search {
g,
order: &order,
s: vec![0i8; n],
lambda: g.h.clone(),
free_h_abs: g.h.iter().map(|v| v.abs()).sum(),
free_abs: g.w.iter().map(|v| v.abs()).sum::<f64>() / 2.0,
nodes: 0,
pruned: 0,
max_nodes: p.max_nodes,
slack,
best,
best_state: best_state.clone(),
hit_limit: false,
gauge_fixed,
undo: Vec::with_capacity(g.nbr.len() + n),
};
st.descend(0, 0.0);
best = st.g.energy(&st.best_state);
best_state = st.best_state;
Outcome {
state: best_state,
energy: best,
proved_optimal: !st.hit_limit,
nodes: st.nodes,
pruned: st.pruned,
hit_limit: st.hit_limit,
slack,
}
}
struct Search<'a> {
g: &'a Graph,
order: &'a [usize],
s: Vec<i8>,
lambda: Vec<f64>,
free_h_abs: f64,
free_abs: f64,
nodes: u64,
pruned: u64,
max_nodes: u64,
slack: f64,
best: f64,
best_state: Vec<i8>,
hit_limit: bool,
gauge_fixed: bool,
undo: Vec<(usize, f64)>,
}
impl Search<'_> {
fn descend(&mut self, depth: usize, fixed_energy: f64) {
if self.hit_limit {
return;
}
self.nodes += 1;
if self.nodes > self.max_nodes {
self.hit_limit = true;
return;
}
if depth == self.order.len() {
let e = self.g.energy(&self.s);
if e < self.best {
self.best = e;
self.best_state.copy_from_slice(&self.s);
}
return;
}
let lb = fixed_energy - self.free_h_abs - self.free_abs;
if lb - self.slack >= self.best {
self.pruned += 1;
return;
}
let i = self.order[depth];
let first: i8 = if self.lambda[i] >= 0.0 { 1 } else { -1 };
let branches: usize = if self.gauge_fixed && depth == 0 { 1 } else { 2 };
for b in 0..branches {
let v = if b == 0 { first } else { -first };
self.fix(i, v, depth, fixed_energy);
if self.hit_limit {
return;
}
}
}
fn fix(&mut self, i: usize, v: i8, depth: usize, fixed_energy: f64) {
let fv = v as f64;
let saved_h_abs = self.free_h_abs;
let saved_abs = self.free_abs;
let new_fixed = fixed_energy - fv * self.lambda[i];
self.s[i] = v;
self.free_h_abs -= self.lambda[i].abs();
let mark = self.undo.len();
let (lo, hi) = (self.g.offset[i], self.g.offset[i + 1]);
for k in lo..hi {
let j = self.g.nbr[k] as usize;
if self.s[j] != 0 {
continue; }
let w = self.g.w[k];
self.free_abs -= w.abs();
self.undo.push((j, self.lambda[j]));
self.free_h_abs -= self.lambda[j].abs();
self.lambda[j] += w * fv;
self.free_h_abs += self.lambda[j].abs();
}
self.descend(depth + 1, new_fixed);
while self.undo.len() > mark {
let (j, old) = self.undo.pop().expect("mark is below len");
self.lambda[j] = old;
}
self.free_abs = saved_abs;
self.free_h_abs = saved_h_abs;
self.s[i] = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::GraphBuilder;
use crate::rng::Pcg;
fn random_graph(n: usize, p: f64, seed: u64, fields: bool) -> Graph {
let mut rng = Pcg::new(seed, 0xB4B4);
let mut gb = GraphBuilder::new(n);
for i in 0..n {
if fields {
gb.bias(i, rng.f64() * 2.0 - 1.0);
}
for j in (i + 1)..n {
if rng.f64() < p {
gb.couple(i, j, rng.f64() * 2.0 - 1.0);
}
}
}
gb.build()
}
fn brute_min(g: &Graph) -> f64 {
let n = g.n;
let mut s = vec![1i8; n];
let mut min = f64::INFINITY;
for mask in 0..(1u64 << n) {
for i in 0..n {
s[i] = if mask >> i & 1 == 1 { 1 } else { -1 };
}
min = min.min(g.energy(&s));
}
min
}
#[test]
fn it_finds_the_true_minimum_and_says_it_proved_it() {
for seed in 0..8u64 {
let fields = seed % 2 == 0;
let g = random_graph(13, 0.4, seed, fields);
let o = solve(&g, &Params::default());
let min = brute_min(&g);
assert!(o.proved_optimal && !o.hit_limit, "seed {seed}: no proof");
assert!(
(o.energy - min).abs() < 1e-9,
"seed {seed} (fields {fields}): branch-and-bound {:.9}, enumeration {min:.9}",
o.energy
);
assert_eq!(o.energy, g.energy(&o.state), "energy must match the state returned");
}
}
#[test]
fn pinning_the_gauge_removes_a_mirror_and_nothing_else() {
let g = random_graph(13, 0.4, 3, false);
let with_gauge = solve(&g, &Params::default());
let tiny = 1e-12;
let mut gb = GraphBuilder::new(g.n);
for i in 0..g.n {
for k in g.offset[i]..g.offset[i + 1] {
let j = g.nbr[k] as usize;
if j > i {
gb.couple(i, j, g.w[k]);
}
}
}
gb.bias(0, tiny);
let g2 = gb.build();
let without = solve(&g2, &Params::default());
assert!(with_gauge.proved_optimal && without.proved_optimal);
assert!(
(with_gauge.energy - without.energy).abs() < 1e-6,
"gauge-pinned {:.9} vs full tree {:.9}",
with_gauge.energy,
without.energy
);
assert!(
with_gauge.nodes < without.nodes,
"pinning the gauge visited {} nodes, the full tree {} -- it should be fewer",
with_gauge.nodes,
without.nodes
);
}
#[test]
fn a_search_that_runs_out_of_budget_does_not_claim_a_proof() {
let g = random_graph(40, 0.3, 5, true);
let o = solve(&g, &Params { max_nodes: 500, incumbent: None });
assert!(o.hit_limit, "500 nodes should not exhaust a 40-spin tree");
assert!(!o.proved_optimal);
assert!(o.nodes <= 501, "nodes {}", o.nodes);
assert_eq!(o.energy, g.energy(&o.state));
assert_eq!(o.state.len(), g.n);
}
#[test]
fn the_bound_prunes_most_of_the_tree() {
let g = random_graph(18, 0.35, 2, true);
let o = solve(&g, &Params::default());
assert!(o.proved_optimal);
let full = 1u64 << 18;
assert!(o.nodes < full, "visited {} nodes of a {full}-leaf tree", o.nodes);
assert!(o.pruned > 0, "no branch was ever cut off, so the bound did nothing");
}
#[test]
fn an_incumbent_changes_the_cost_and_not_the_result() {
for seed in 0..4u64 {
let g = random_graph(12, 0.4, 40 + seed, seed % 2 == 0);
let plain = solve(&g, &Params::default());
let seeded = solve(&g, &Params { incumbent: Some(plain.state.clone()), ..Params::default() });
assert!(seeded.proved_optimal);
assert!(
(seeded.energy - plain.energy).abs() < 1e-12,
"seed {seed}: {:.9} vs {:.9}",
seeded.energy,
plain.energy
);
assert!(seeded.nodes <= plain.nodes, "a known-optimal incumbent should not cost nodes");
let bad = solve(&g, &Params { incumbent: Some(vec![0i8; g.n]), ..Params::default() });
assert!((bad.energy - plain.energy).abs() < 1e-12);
let short = solve(&g, &Params { incumbent: Some(vec![1i8; g.n - 1]), ..Params::default() });
assert!((short.energy - plain.energy).abs() < 1e-12);
}
}
#[test]
fn an_empty_graph_is_proved_immediately() {
let g = GraphBuilder::new(0).build();
let o = solve(&g, &Params::default());
assert!(o.proved_optimal && o.nodes == 0 && o.energy == 0.0 && o.state.is_empty());
}
#[test]
fn the_prune_slack_is_positive_and_negligible() {
let g = random_graph(14, 0.4, 6, true);
let o = solve(&g, &Params::default());
let scale: f64 = g.w.iter().map(|v| v.abs()).sum::<f64>();
assert!(o.slack > 0.0, "a zero slack cannot absorb the accumulated rounding");
assert!(o.slack < scale * 1e-9, "slack {} against a weight scale of {scale}", o.slack);
}
}