use std::time::Instant;
use super::elim::ElimYield;
use crate::cnf::VarId;
use crate::cnf::{Clause, CnfFormula, Literal};
use crate::diagnostics::diag;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(crate) enum FrozenEquiv {
Ignore,
ForceShowRep,
}
#[cfg(test)]
pub(crate) fn post_dve_strengthen(
dve: &mut super::types::DveResult,
frozen: &rustc_hash::FxHashSet<VarId>,
) {
let mut meter =
crate::preprocess::meter::PreprocessMeter::new(crate::config::PreprocessClock::WallClock);
post_dve_strengthen_with_meter(dve, frozen, &mut meter)
}
pub(crate) fn post_dve_strengthen_with_meter(
dve: &mut super::types::DveResult,
frozen: &rustc_hash::FxHashSet<VarId>,
meter: &mut crate::preprocess::meter::PreprocessMeter,
) {
if dve.formula.clauses.is_empty() || dve.formula.num_vars == 0 {
return;
}
let clauses_before = dve.formula.clauses.len();
let vars_before = dve.formula.num_vars;
let inner_frozen: rustc_hash::FxHashSet<VarId> = if frozen.is_empty() {
rustc_hash::FxHashSet::default()
} else {
(0..vars_before as usize)
.filter(|&p| {
let din = match &dve.renumbering {
Some(r) => r.old_id(VarId(p as u32)).idx(),
None => p,
};
frozen.contains(&VarId(din as u32))
})
.map(|p| VarId(p as u32))
.collect()
};
let mapping = super::super::gates::detect_gates(&dve.formula);
let known_defined = mapping.eliminated;
let inner = super::pipeline::preprocess_dve_with_meter(
&dve.formula,
super::pipeline::DveConfig {
max_rounds: 10,
time_limit_ms: 2_000,
keep_original_vars: false,
known_defined: &known_defined,
frozen: &inner_frozen,
frozen_equiv: FrozenEquiv::Ignore,
},
meter,
);
dve.elapsed_ms = dve.elapsed_ms.saturating_add(inner.elapsed_ms);
if inner.total_eliminated() == 0 {
return;
}
let old_renumbering = dve.renumbering.take();
dve.renumbering = match (old_renumbering.as_ref(), inner.renumbering.as_ref()) {
(Some(old), Some(inner_r)) => Some(old.compose(inner_r)),
(Some(_), None) => old_renumbering.clone(),
(None, inner_r) => inner_r.cloned(),
};
let to_dve_input = |p: u32| match old_renumbering.as_ref() {
Some(r) => r.old_id(VarId(p)).0,
None => p,
};
for (p, &fate) in inner.fates.iter().enumerate() {
if !fate.eliminated() {
continue;
}
dve.fates[to_dve_input(p as u32) as usize] = match fate {
super::types::DveFate::Equiv { rep } => super::types::DveFate::Equiv {
rep: Literal::new(VarId(to_dve_input(rep.var.0)), rep.positive),
},
other => other,
};
}
let (inner_defined, inner_equiv, inner_free) =
(inner.num_defined(), inner.num_equiv(), inner.num_free());
dve.formula = inner.formula;
dve.definition_clauses.extend(inner.definition_clauses);
diag!(
"[post-dve] {} → {} vars, {} → {} clauses ({} defined, {} equiv, {} free)",
vars_before,
dve.formula.num_vars,
clauses_before,
dve.formula.clauses.len(),
inner_defined,
inner_equiv,
inner_free,
);
}
pub(super) struct EquivState<'a> {
pub(super) fates: &'a mut [super::types::DveFate],
pub(super) representative: &'a mut [i32],
}
pub(super) fn merge_equivalences(
clauses: &mut Vec<Clause>,
num_vars: usize,
state: &mut EquivState<'_>,
frozen: &rustc_hash::FxHashSet<VarId>,
policy: FrozenEquiv,
) -> ElimYield {
let num_lits = num_vars * 2;
let mut adj: Vec<Vec<u32>> = vec![Vec::new(); num_lits];
for clause in clauses.iter() {
if clause.literals.len() == 2 {
let (l0, l1) = (&clause.literals[0], &clause.literals[1]);
adj[lit_to_idx(l0.var.0 as usize, !l0.positive)]
.push(lit_to_idx(l1.var.0 as usize, l1.positive) as u32);
adj[lit_to_idx(l1.var.0 as usize, !l1.positive)]
.push(lit_to_idx(l0.var.0 as usize, l0.positive) as u32);
}
}
let sccs = tarjan_scc(&adj, num_lits);
let mut scc_id = vec![0u32; num_lits];
for (id, scc) in sccs.iter().enumerate() {
for &node in scc {
scc_id[node as usize] = id as u32;
}
}
let mut equiv_count = 0usize;
let mut equiv_def_clauses: Vec<Vec<Clause>> = Vec::new();
let mut rep_map: Vec<Literal> = (0..num_vars)
.map(|v| Literal::pos(VarId(v as u32)))
.collect();
for scc in &sccs {
if scc.len() <= 1 {
continue;
}
for &node in scc {
let var = node as usize / 2;
let pos = (node as usize).is_multiple_of(2);
let neg_idx = lit_to_idx(var, !pos);
if scc_id[node as usize] == scc_id[neg_idx] {
clauses.clear();
clauses.push(Clause::new(vec![]));
return ElimYield {
eliminated: 0,
definitions: Vec::new(),
};
}
}
let force_show_rep = policy == FrozenEquiv::ForceShowRep
&& !frozen.is_empty()
&& scc
.iter()
.any(|&node| frozen.contains(&VarId((node as usize / 2) as u32)));
let mut rep_var = u32::MAX;
let mut rep_positive = true;
for &node in scc {
let var = node as usize / 2;
let positive = (node as usize).is_multiple_of(2);
if state.fates[var].eliminated() {
continue;
}
if force_show_rep && !frozen.contains(&VarId(var as u32)) {
continue;
}
if (var as u32) < rep_var {
rep_var = var as u32;
rep_positive = positive;
}
}
if rep_var == u32::MAX {
continue; }
for &node in scc {
let var = node as usize / 2;
let positive = (node as usize).is_multiple_of(2);
if var as u32 == rep_var || state.fates[var].eliminated() {
continue;
}
let same_pol = positive == rep_positive;
let rep_lit = Literal::new(VarId(rep_var), same_pol);
rep_map[var] = rep_lit;
state.fates[var] = super::types::DveFate::Equiv { rep: rep_lit };
state.representative[var] = if same_pol {
rep_var as i32
} else {
-(rep_var as i32)
};
let v_id = VarId(var as u32);
let rep_id = VarId(rep_var);
let def = if same_pol {
vec![
Clause::new(vec![Literal::pos(v_id), Literal::neg(rep_id)]),
Clause::new(vec![Literal::neg(v_id), Literal::pos(rep_id)]),
]
} else {
vec![
Clause::new(vec![Literal::pos(v_id), Literal::pos(rep_id)]),
Clause::new(vec![Literal::neg(v_id), Literal::neg(rep_id)]),
]
};
equiv_def_clauses.push(def);
equiv_count += 1;
}
}
if equiv_count == 0 {
return ElimYield {
eliminated: 0,
definitions: Vec::new(),
};
}
let mut new_clauses: Vec<Clause> = Vec::with_capacity(clauses.len());
for clause in clauses.iter() {
let mut new_lits: Vec<Literal> = clause
.literals
.iter()
.map(|lit| {
let rep = rep_map[lit.var.0 as usize];
if lit.positive { rep } else { rep.negated() }
})
.collect();
new_lits.sort_by_key(|l| (l.var, !l.positive));
new_lits.dedup();
let is_tautology = new_lits.windows(2).any(|w| w[0].var == w[1].var);
if !is_tautology {
new_clauses.push(Clause::new(new_lits));
}
}
*clauses = new_clauses;
super::elim::dedup_clauses(clauses);
ElimYield {
eliminated: equiv_count,
definitions: equiv_def_clauses,
}
}
#[cfg(test)]
pub(super) fn strengthen_clauses(
clauses: &mut Vec<Clause>,
num_vars: usize,
stage_deadline: Option<Instant>,
) -> bool {
let mut meter =
crate::preprocess::meter::PreprocessMeter::new(crate::config::PreprocessClock::WallClock);
strengthen_clauses_with_meter(clauses, num_vars, stage_deadline, &mut meter)
}
pub(super) fn strengthen_clauses_with_meter(
clauses: &mut Vec<Clause>,
num_vars: usize,
stage_deadline: Option<Instant>,
meter: &mut crate::preprocess::meter::PreprocessMeter,
) -> bool {
if clauses.is_empty() {
return false;
}
let len_before = clauses.len();
let total_lits_before: usize = clauses.iter().map(|c| c.literals.len()).sum();
let formula = CnfFormula {
num_vars: num_vars as u32,
clauses: std::mem::take(clauses),
};
let deadline = stage_deadline.map(|d| {
let now = Instant::now();
now + d.saturating_duration_since(now) / 2
});
let (strengthened, _forced) =
super::super::cadical::preprocess_cadical_with_meter(&formula, 1, deadline, meter);
if strengthened.clauses.len() == len_before {
let total_lits_after: usize = strengthened.clauses.iter().map(|c| c.literals.len()).sum();
if total_lits_before == total_lits_after {
*clauses = formula.clauses;
return false;
}
}
*clauses = strengthened.clauses;
true
}
fn tarjan_scc(adj: &[Vec<u32>], n: usize) -> Vec<Vec<u32>> {
let adj_usize: Vec<Vec<usize>> = adj[..n]
.iter()
.map(|neighbors| neighbors.iter().map(|&w| w as usize).collect())
.collect();
let groups_usize = super::super::tarjan::tarjan_scc_groups(n, &adj_usize);
groups_usize
.into_iter()
.map(|group| group.into_iter().map(|v| v as u32).collect())
.collect()
}
use crate::cnf::occ::literal_index as lit_to_idx;