use std::collections::HashSet;
use std::time::{Duration, Instant};
use super::cadical_ffi::{Bounded, CaDiCal, Status};
use crate::cnf::{CnfFormula, Literal, VarId};
use super::backbone::{BackboneResult, EquivResult, read_model, refine_candidates};
use super::cadical::WallClockTerminator;
use super::equivalence::EquivMapping;
use crate::bundle::PreprocessPhase;
use crate::cnf::occ;
enum ProbeSolver<'a> {
Direct(&'a mut CaDiCal),
Wall(Bounded<'a, WallClockTerminator>),
}
impl std::ops::Deref for ProbeSolver<'_> {
type Target = CaDiCal;
fn deref(&self) -> &CaDiCal {
match self {
ProbeSolver::Direct(solver) => solver,
ProbeSolver::Wall(solver) => solver,
}
}
}
impl std::ops::DerefMut for ProbeSolver<'_> {
fn deref_mut(&mut self) -> &mut CaDiCal {
match self {
ProbeSolver::Direct(solver) => solver,
ProbeSolver::Wall(solver) => solver,
}
}
}
const MAX_CONFLICTS: i32 = 64_000;
#[inline]
fn lit_true_in_model(lit: i32, model: &[i32]) -> bool {
let mv = model[VarId::from_dimacs(lit).idx()];
(lit > 0 && mv > 0) || (lit < 0 && mv < 0)
}
fn map_lit(d: i32, mapping: &Option<EquivMapping>) -> Literal {
let lit = Literal::from(d);
match mapping {
None => lit,
Some(m) => {
let rep = m.var_to_rep[lit.var.idx()];
if lit.positive { rep } else { rep.negated() }
}
}
}
pub(super) struct ProbeEngine {
solver: CaDiCal,
num_vars: usize,
freq: Vec<u32>,
pub(super) partition: Partition,
}
pub(super) struct Partition {
pub(super) classes: Vec<Vec<i32>>,
pub(super) top: usize,
seeded: bool,
pub(super) confirmed_backbone: Vec<Literal>,
confirmed_equiv: Vec<(i32, i32)>,
}
impl Partition {
pub(super) fn new() -> Self {
Partition {
classes: Vec::new(),
top: usize::MAX,
seeded: false,
confirmed_backbone: Vec::new(),
confirmed_equiv: Vec::new(),
}
}
pub(super) fn observe_model(&mut self, model: &[i32]) {
let top = self.top;
let old = std::mem::take(&mut self.classes);
let mut new_top = usize::MAX;
for (ci, class) in old.into_iter().enumerate() {
let mut t_half = Vec::new();
let mut f_half = Vec::new();
for lit in class {
if lit_true_in_model(lit, model) {
t_half.push(lit);
} else {
f_half.push(lit);
}
}
if ci == top {
new_top = self.classes.len();
self.classes.push(t_half);
if f_half.len() >= 2 {
self.classes.push(f_half);
}
} else {
if t_half.len() >= 2 {
self.classes.push(t_half);
}
if f_half.len() >= 2 {
self.classes.push(f_half);
}
}
}
self.top = new_top;
}
fn remove_lits(&mut self, lits: &HashSet<i32>) {
for class in &mut self.classes {
class.retain(|l| !lits.contains(l));
}
}
}
impl ProbeEngine {
pub(super) fn new(formula: &CnfFormula) -> Option<Self> {
let mut solver = CaDiCal::new()?;
for clause in &formula.clauses {
for lit in &clause.literals {
solver.add(lit.to_dimacs());
}
solver.add(0);
}
solver.limit(c"conflicts", 1_000_000);
let num_vars = formula.num_vars as usize;
let freq = occ::frequency(&formula.clauses, num_vars);
Some(ProbeEngine {
solver,
num_vars,
freq,
partition: Partition::new(),
})
}
#[cfg(test)]
pub(super) fn run_backbone(&mut self, budget: Duration) -> BackboneResult {
let mut meter =
super::meter::PreprocessMeter::new(crate::config::PreprocessClock::WallClock);
self.run_backbone_with_meter(budget, &mut meter)
}
pub(super) fn run_backbone_with_meter(
&mut self,
budget: Duration,
meter: &mut super::meter::PreprocessMeter,
) -> BackboneResult {
let start = Instant::now();
let mark = meter.begin(PreprocessPhase::Backbone, budget);
let nv = self.num_vars;
let empty = |solve_ms: u64, unsat: bool| BackboneResult {
forced: Vec::new(),
probes_completed: 0,
solve_ms,
unsat,
fixed_found: 0,
flippable_eliminated: 0,
model_eliminated: 0,
elapsed_ms: start.elapsed().as_millis() as u64,
};
if nv == 0 {
let result = empty(0, false);
meter.finish_phase(mark);
return result;
}
let mut solver = if meter.deterministic() {
ProbeSolver::Direct(&mut self.solver)
} else {
ProbeSolver::Wall(Bounded::new(
&mut self.solver,
WallClockTerminator::new(budget),
))
};
let status = meter.solve(PreprocessPhase::Backbone, &mut solver);
let solve_ms = start.elapsed().as_millis() as u64;
match status {
Status::Unsatisfiable => {
let result = empty(solve_ms, true);
meter.finish_phase(mark);
return result;
}
Status::Satisfiable => {}
_ => {
let result = empty(solve_ms, false);
meter.finish_phase(mark);
return result;
}
}
let model = read_model(&mut solver, nv);
let mut top_class = Vec::new();
for &lit in &model {
if lit != 0 {
top_class.push(lit);
}
}
self.partition.classes = vec![top_class];
self.partition.top = 0;
self.partition.seeded = true;
let (fixed_found, flippable_eliminated) =
harvest_fixed_and_flippable(&mut self.partition, &mut solver, &model, nv);
let mut candidates: Vec<i32> = self.partition.classes[self.partition.top].clone();
candidates.sort_unstable_by(|&a, &b| {
let fa = self.freq[VarId::from_dimacs(a).idx()];
let fb = self.freq[VarId::from_dimacs(b).idx()];
fb.cmp(&fa)
});
let probed = probe_loop(
&mut self.partition,
&mut solver,
&mut candidates,
nv,
mark,
budget,
meter,
);
let recovered = recover_deferred(
&mut self.partition,
&mut solver,
&probed.deferred,
nv,
mark,
budget,
meter,
);
let result = BackboneResult {
forced: self.partition.confirmed_backbone.clone(),
probes_completed: probed.probes_completed + recovered,
solve_ms,
unsat: false,
fixed_found,
flippable_eliminated,
model_eliminated: probed.model_eliminated,
elapsed_ms: start.elapsed().as_millis() as u64,
};
meter.finish_phase(mark);
result
}
pub(super) fn ingest_tarjan_equivs(&mut self, mapping: &EquivMapping) {
let mut drop: HashSet<i32> = HashSet::new();
for (v, &rep) in mapping.var_to_rep.iter().enumerate() {
if rep.var.0 != v as u32 {
let d = (v as i32) + 1;
drop.insert(d);
drop.insert(-d);
}
}
if !drop.is_empty() {
self.partition.remove_lits(&drop);
}
}
#[cfg(test)]
pub(super) fn run_equiv(
&mut self,
budget: Duration,
mapping2: &Option<EquivMapping>,
) -> EquivResult {
let mut meter =
super::meter::PreprocessMeter::new(crate::config::PreprocessClock::WallClock);
self.run_equiv_with_meter(budget, mapping2, &mut meter)
}
pub(super) fn run_equiv_with_meter(
&mut self,
budget: Duration,
mapping2: &Option<EquivMapping>,
meter: &mut super::meter::PreprocessMeter,
) -> EquivResult {
let start = Instant::now();
let mark = meter.begin(PreprocessPhase::Equivalence, budget);
let nv = self.num_vars;
if !self.partition.seeded {
let result = EquivResult {
equivalences: Vec::new(),
probes_completed: 0,
unsat: false,
elapsed_ms: start.elapsed().as_millis() as u64,
};
meter.finish_phase(mark);
return result;
}
self.partition.top = usize::MAX;
let mut solver = if meter.deterministic() {
ProbeSolver::Direct(&mut self.solver)
} else {
ProbeSolver::Wall(Bounded::new(
&mut self.solver,
WallClockTerminator::new(budget),
))
};
let mut probes_completed = 0;
loop {
if meter.elapsed(mark) >= budget {
break;
}
self.partition
.classes
.sort_unstable_by_key(|c| std::cmp::Reverse(c.len()));
while self.partition.classes.last().is_some_and(|p| p.len() < 2) {
self.partition.classes.pop();
}
if self.partition.classes.is_empty() || self.partition.classes[0].len() < 2 {
break;
}
let class = self.partition.classes.remove(0);
let rep = class[0];
let mut confirmed = vec![rep];
let mut remaining: Vec<i32> = class[1..].to_vec();
let mut i = 0;
while i < remaining.len() {
if meter.elapsed(mark) >= budget {
break;
}
let candidate = remaining[i];
probes_completed += 1;
solver.assume(rep);
solver.assume(-candidate);
if let Some(cap) = meter.equivalence_conflict_cap() {
solver.limit(c"conflicts", cap);
}
match meter.solve(PreprocessPhase::Equivalence, &mut solver) {
Status::Satisfiable => {
let new_model = read_model(&mut solver, nv);
self.partition.observe_model(&new_model);
let (stay, split) = refine_candidates(&remaining[i..], &new_model, true);
remaining.truncate(i);
remaining.extend(stay);
if split.len() >= 2 {
self.partition.classes.push(split);
}
continue; }
Status::Unsatisfiable => {} _ => {
i += 1;
continue;
}
}
probes_completed += 1;
solver.assume(-rep);
solver.assume(candidate);
if let Some(cap) = meter.equivalence_conflict_cap() {
solver.limit(c"conflicts", cap);
}
match meter.solve(PreprocessPhase::Equivalence, &mut solver) {
Status::Unsatisfiable => {
confirmed.push(candidate);
remaining.remove(i);
}
Status::Satisfiable => {
let new_model = read_model(&mut solver, nv);
self.partition.observe_model(&new_model);
let (stay, split) = refine_candidates(&remaining[i..], &new_model, false);
remaining.truncate(i);
remaining.extend(stay);
if split.len() >= 2 {
self.partition.classes.push(split);
}
continue; }
_ => {
i += 1;
}
}
}
if confirmed.len() >= 2 {
for &other in &confirmed[1..] {
self.partition.confirmed_equiv.push((rep, other));
}
}
if remaining.len() >= 2 {
self.partition.classes.push(remaining);
}
}
let mut equivalences = Vec::new();
for &(a, b) in &self.partition.confirmed_equiv {
let la = map_lit(a, mapping2);
let lb = map_lit(b, mapping2);
if la.var == lb.var {
continue;
}
equivalences.push((la, lb));
}
let result = EquivResult {
equivalences,
probes_completed,
unsat: false,
elapsed_ms: start.elapsed().as_millis() as u64,
};
meter.finish_phase(mark);
result
}
}
fn confirm_backbone(
partition: &mut Partition,
solver: &mut ProbeSolver<'_>,
lits: impl IntoIterator<Item = i32>,
) {
let mut confirmed: HashSet<i32> = HashSet::new();
for lit in lits {
partition.confirmed_backbone.push(Literal::from(lit));
solver.add(lit);
solver.add(0);
confirmed.insert(lit);
}
partition.remove_lits(&confirmed);
}
fn harvest_fixed_and_flippable(
partition: &mut Partition,
solver: &mut ProbeSolver<'_>,
model: &[i32],
nv: usize,
) -> (usize, usize) {
let mut fixed_found = 0;
let mut flippable_eliminated = 0;
let mut remove: HashSet<i32> = HashSet::new();
for i in 0..nv {
let dimacs = VarId(i as u32).to_dimacs();
let f = solver.fixed(dimacs);
if f != 0 {
let lit = if f > 0 { dimacs } else { -dimacs };
partition.confirmed_backbone.push(Literal::from(lit));
fixed_found += 1;
remove.insert(lit);
}
}
for (i, &val) in model.iter().enumerate() {
if val == 0 {
continue;
}
if solver.fixed(VarId(i as u32).to_dimacs()) != 0 {
continue; }
if solver.flippable(-val) {
flippable_eliminated += 1;
remove.insert(val);
}
}
if !remove.is_empty() {
partition.remove_lits(&remove);
}
(fixed_found, flippable_eliminated)
}
struct ProbeRun {
probes_completed: usize,
model_eliminated: usize,
deferred: Vec<i32>,
}
fn probe_loop(
partition: &mut Partition,
solver: &mut ProbeSolver<'_>,
candidates: &mut Vec<i32>,
nv: usize,
mark: super::meter::PhaseMark,
budget: Duration,
meter: &mut super::meter::PreprocessMeter,
) -> ProbeRun {
let mut probes_completed = 0;
let mut model_eliminated = 0;
let mut chunk_limit: usize = 1;
let mut pos = 0;
let mut deferred: Vec<i32> = Vec::new();
while pos < candidates.len() {
if meter.elapsed(mark) >= budget {
break;
}
let remaining = candidates.len() - pos;
let chunk_size = chunk_limit.min(remaining);
probes_completed += 1;
let probe = if chunk_size == 1 {
solver.limit(c"conflicts", MAX_CONFLICTS);
solver.assume(-candidates[pos]);
meter.solve(PreprocessPhase::Backbone, solver)
} else {
solver.limit(c"conflicts", 1_000_000);
for &cand in &candidates[pos..pos + chunk_size] {
solver.constrain(-cand);
}
solver.constrain(0);
meter.solve(PreprocessPhase::Backbone, solver)
};
match probe {
Status::Unsatisfiable => {
confirm_backbone(
partition,
solver,
candidates[pos..pos + chunk_size].iter().copied(),
);
pos += chunk_size;
chunk_limit = chunk_limit.saturating_mul(8).max(1);
}
Status::Satisfiable => {
let new_model = read_model(solver, nv);
partition.observe_model(&new_model);
let top_set: HashSet<i32> =
partition.classes[partition.top].iter().copied().collect();
let mut write = pos;
for read in pos..candidates.len() {
if top_set.contains(&candidates[read]) {
candidates[write] = candidates[read];
write += 1;
} else {
model_eliminated += 1;
}
}
candidates.truncate(write);
chunk_limit = 1;
}
_ => {
if meter.elapsed(mark) < budget {
deferred.extend(candidates.drain(pos..pos + chunk_size));
chunk_limit = 1;
} else {
break;
}
}
}
}
ProbeRun {
probes_completed,
model_eliminated,
deferred,
}
}
fn recover_deferred(
partition: &mut Partition,
solver: &mut ProbeSolver<'_>,
deferred: &[i32],
nv: usize,
mark: super::meter::PhaseMark,
budget: Duration,
meter: &mut super::meter::PreprocessMeter,
) -> usize {
const RECOVERY_CAP: i32 = 1_000;
let mut probes_completed = 0;
for &lit in deferred {
if meter.elapsed(mark) >= budget {
break;
}
probes_completed += 1;
solver.limit(c"conflicts", RECOVERY_CAP);
solver.assume(-lit);
match meter.solve(PreprocessPhase::Backbone, solver) {
Status::Unsatisfiable => confirm_backbone(partition, solver, [lit]),
Status::Satisfiable => {
let new_model = read_model(solver, nv);
partition.observe_model(&new_model);
}
_ => {}
}
}
probes_completed
}