use crate::cnf::VarId;
use crate::cnf::{Clause, Literal};
use crate::cnf::occ;
use super::definability::{
MAX_DUAL_CNF_CLAUSES, PRIMAL_GRAPH_MAX_VARS, PrimalGraph, is_ve_candidate,
pick_def_vars_with_meter,
};
use super::types::DveFate;
struct PolaritySplit {
pos: Vec<Clause>,
neg: Vec<Clause>,
remaining: Vec<Clause>,
originals: Vec<Clause>,
}
fn split_on(clauses: &mut Vec<Clause>, v: u32) -> PolaritySplit {
let mut split = PolaritySplit {
pos: Vec::new(),
neg: Vec::new(),
remaining: Vec::new(),
originals: Vec::new(),
};
for clause in clauses.drain(..) {
let mut found_pos = false;
let mut found_neg = false;
for lit in &clause.literals {
if lit.var.0 == v {
if lit.positive {
found_pos = true;
} else {
found_neg = true;
}
}
}
if found_pos || found_neg {
split.originals.push(clause.clone());
let stripped: Vec<Literal> = clause
.literals
.iter()
.filter(|l| l.var.0 != v)
.copied()
.collect();
let target = if found_pos {
&mut split.pos
} else {
&mut split.neg
};
target.push(Clause::new(stripped));
} else {
split.remaining.push(clause);
}
}
split
}
pub(super) fn elim_vars(
clauses: &mut Vec<Clause>,
vars_to_elim: &[u32],
max_clauses: usize,
frozen: &rustc_hash::FxHashSet<VarId>,
) -> (Vec<u32>, Vec<Literal>, Vec<Vec<Clause>>) {
let mut eliminated_ids: Vec<u32> = Vec::new();
let mut forced_lits = Vec::new();
let mut all_def_clauses: Vec<Vec<Clause>> = Vec::new();
for &v in vars_to_elim {
if clauses.len() > max_clauses {
break;
}
let PolaritySplit {
pos: pos_clauses,
neg: neg_clauses,
mut remaining,
originals: original_clauses_for_v,
} = split_on(clauses, v);
if pos_clauses.is_empty() && neg_clauses.is_empty() {
*clauses = remaining;
continue;
}
let mut resolvents: Vec<Clause> = Vec::new();
let mut abort = false;
for c1 in &pos_clauses {
for c2 in &neg_clauses {
let mut merged: Vec<Literal> = Vec::new();
let mut is_tautology = false;
let mut i = 0;
let mut j = 0;
while i < c1.literals.len() && j < c2.literals.len() {
let l1 = &c1.literals[i];
let l2 = &c2.literals[j];
if l1.var < l2.var {
merged.push(*l1);
i += 1;
} else if l1.var > l2.var {
merged.push(*l2);
j += 1;
} else {
if l1.positive == l2.positive {
merged.push(*l1);
i += 1;
j += 1;
} else {
is_tautology = true;
break;
}
}
}
if !is_tautology {
merged.extend_from_slice(&c1.literals[i..]);
merged.extend_from_slice(&c2.literals[j..]);
if merged.len() <= 1 {
if let Some(&lit) = merged.first() {
if frozen.contains(&lit.var) {
resolvents.push(Clause::new(merged));
} else {
forced_lits.push(lit);
}
} else {
resolvents.push(Clause::new(merged));
}
} else {
resolvents.push(Clause::new(merged));
}
}
if remaining.len() + resolvents.len() > max_clauses {
abort = true;
break;
}
}
if abort {
break;
}
}
if abort {
for mut c in pos_clauses {
c.literals.push(Literal::pos(VarId(v)));
c.literals.sort_by_key(|l| l.var);
remaining.push(c);
}
for mut c in neg_clauses {
c.literals.push(Literal::neg(VarId(v)));
c.literals.sort_by_key(|l| l.var);
remaining.push(c);
}
*clauses = remaining;
break;
}
if pos_clauses.is_empty() || neg_clauses.is_empty() {
remaining.extend(original_clauses_for_v);
*clauses = remaining;
continue;
}
remaining.extend(resolvents);
*clauses = remaining;
eliminated_ids.push(v);
all_def_clauses.push(original_clauses_for_v);
}
(eliminated_ids, forced_lits, all_def_clauses)
}
pub(super) struct ElimYield {
pub(super) eliminated: usize,
pub(super) definitions: Vec<Vec<Clause>>,
}
pub(super) struct RoundStats {
pub(super) round: usize,
pub(super) dve_elim: usize,
pub(super) equiv_elim: usize,
pub(super) round1_dve_elim: usize,
pub(super) vars_before: usize,
pub(super) vars_after: usize,
pub(super) clauses_before: usize,
pub(super) clauses_after: usize,
pub(super) orig_clause_count: usize,
pub(super) progress: bool,
}
pub(super) fn should_terminate_dve(s: &RoundStats) -> bool {
if s.round == 0 && s.dve_elim > 0 && s.equiv_elim == 0 {
let elim_rate = s.dve_elim as f64 / s.vars_before.max(1) as f64;
if elim_rate < 0.02 {
return true;
}
}
if !s.progress {
return true;
}
if s.clauses_after > (s.orig_clause_count as f64 * 1.1) as usize {
return true;
}
if s.vars_after == s.vars_before && s.clauses_after == s.clauses_before {
return true;
}
let dim_threshold = (s.round1_dve_elim / 20).max(3);
if s.dve_elim < dim_threshold && s.equiv_elim == 0 && s.round >= 2 {
return true;
}
false
}
pub(super) fn count_active_vars(clauses: &[Clause]) -> usize {
clauses
.iter()
.flat_map(|c| c.literals.iter().map(|l| l.var.0))
.collect::<std::collections::HashSet<_>>()
.len()
}
pub(super) fn sort_clause_literals(clauses: &mut [Clause]) {
for clause in clauses.iter_mut() {
clause.literals.sort_by_key(|l| l.var);
}
}
pub(super) fn dedup_clauses(clauses: &mut Vec<Clause>) {
for clause in clauses.iter_mut() {
clause.literals.sort_by_key(|l| (l.var, !l.positive));
clause.literals.dedup();
}
clauses.sort_by(|a, b| {
a.literals.len().cmp(&b.literals.len()).then_with(|| {
a.literals
.iter()
.zip(b.literals.iter())
.map(|(la, lb)| la.var.cmp(&lb.var).then(la.positive.cmp(&lb.positive)))
.find(|o| !o.is_eq())
.unwrap_or(std::cmp::Ordering::Equal)
})
});
clauses.dedup();
}
pub(super) fn propagate_forced(
clauses: &mut Vec<Clause>,
forced: &[Literal],
frozen: &rustc_hash::FxHashSet<VarId>,
) {
if forced.is_empty() {
return;
}
use rustc_hash::FxHashSet;
let mut assigned: FxHashSet<(u32, bool)> = FxHashSet::default();
let mut queue: Vec<Literal> = forced
.iter()
.copied()
.filter(|l| !frozen.contains(&l.var))
.collect();
while let Some(lit) = queue.pop() {
if !assigned.insert((lit.var.0, lit.positive)) {
continue;
}
let mut i = 0;
while i < clauses.len() {
let contains_lit = clauses[i]
.literals
.iter()
.any(|l| l.var == lit.var && l.positive == lit.positive);
if contains_lit {
clauses.swap_remove(i);
continue;
}
let contains_negation = clauses[i]
.literals
.iter()
.any(|l| l.var == lit.var && l.positive != lit.positive);
if contains_negation {
clauses[i].literals.retain(|l| l.var != lit.var);
if clauses[i].literals.is_empty() {
return;
}
if clauses[i].literals.len() == 1 {
let unit = clauses[i].literals[0];
if !frozen.contains(&unit.var) {
queue.push(unit);
}
}
}
i += 1;
}
}
}
pub(super) fn apply_elimination(
clauses: &mut Vec<Clause>,
defined: &[u32],
fates: &mut [DveFate],
max_clauses: usize,
frozen: &rustc_hash::FxHashSet<VarId>,
) -> ElimYield {
sort_clause_literals(clauses);
let (elim_ids, forced, def_clauses) = elim_vars(clauses, defined, max_clauses, frozen);
let elim_count = elim_ids.len();
for v in elim_ids {
fates[v as usize] = DveFate::Defined;
}
if !forced.is_empty() {
propagate_forced(clauses, &forced, frozen);
}
dedup_clauses(clauses);
ElimYield {
eliminated: elim_count,
definitions: def_clauses,
}
}
pub(super) fn dve_round(
clauses: &mut Vec<Clause>,
num_vars: usize,
fates: &mut [DveFate],
time_limit_ms: u64,
known_defined: &rustc_hash::FxHashSet<VarId>,
frozen: &rustc_hash::FxHashSet<VarId>,
meter: &mut crate::preprocess::meter::PreprocessMeter,
) -> ElimYield {
let graph = if num_vars <= PRIMAL_GRAPH_MAX_VARS {
Some(PrimalGraph::new(num_vars, clauses))
} else {
None
};
let freq = occ::literal_frequency(clauses, num_vars);
let mut sat_candidates: Vec<u32> = Vec::new();
let mut preknown: Vec<u32> = Vec::new();
let bypass_ve_filter = graph.is_none();
for v in 0..num_vars {
if fates[v].eliminated() {
continue;
}
if frozen.contains(&VarId(v as u32)) {
continue;
}
if known_defined.contains(&VarId(v as u32)) {
let pf = freq[v * 2] as u64;
let nf = freq[v * 2 + 1] as u64;
if pf == 0 && nf == 0 {
continue;
}
if bypass_ve_filter || is_ve_candidate(graph.as_ref(), &freq, v) {
preknown.push(v as u32);
}
} else if is_ve_candidate(graph.as_ref(), &freq, v) {
sat_candidates.push(v as u32);
}
}
if preknown.is_empty() && sat_candidates.is_empty() {
return ElimYield {
eliminated: 0,
definitions: Vec::new(),
};
}
let freq_key = |&v: &u32| {
let v = v as usize;
freq[v * 2] + freq[v * 2 + 1]
};
sat_candidates.sort_by_key(freq_key);
preknown.sort_by_key(freq_key);
let mut yielded = ElimYield {
eliminated: 0,
definitions: Vec::new(),
};
if !preknown.is_empty() {
let max_clauses = clauses.len();
let step = apply_elimination(clauses, &preknown, fates, max_clauses, frozen);
yielded.eliminated += step.eliminated;
yielded.definitions.extend(step.definitions);
}
let preknown_remaining = known_defined.iter().any(|v| !fates[v.idx()].eliminated());
if !preknown_remaining && !sat_candidates.is_empty() && clauses.len() <= MAX_DUAL_CNF_CLAUSES {
sat_candidates.retain(|&v| !fates[v as usize].eliminated());
if !sat_candidates.is_empty() {
let sat_defined =
pick_def_vars_with_meter(clauses, num_vars, &sat_candidates, time_limit_ms, meter);
if !sat_defined.is_empty() {
let max_clauses = clauses.len();
let step = apply_elimination(clauses, &sat_defined, fates, max_clauses, frozen);
yielded.eliminated += step.eliminated;
yielded.definitions.extend(step.definitions);
}
}
}
yielded
}