use num_traits::One;
use super::stage::{ArjunOutcome, arjun_stage, simplify_outcome};
use super::*;
pub(super) fn count_preserving_bundle_with_stage1(
formula: &CnfFormula,
meta: &CnfMeta,
config: &RunConfig,
mode: Mode,
) -> Result<(PreprocessBundle, CountStage1), VitriError> {
let stage1 = count_stage1(formula, meta, config, mode);
let bundle = finish_count_preserving_attempt(&stage1, config)?;
Ok((bundle, stage1))
}
pub(super) struct CountStage1 {
simplified: SimplifiedFormula,
stage1_weights: Weights<Reduced>,
stage1_lift: BigRational,
stages: StageReport,
telemetry: PreprocessTelemetry,
mode: Mode,
}
pub(super) fn count_stage1(
formula: &CnfFormula,
meta: &CnfMeta,
config: &RunConfig,
mode: Mode,
) -> CountStage1 {
let started = std::time::Instant::now();
let weighted = mode.is_weighted();
let orig_nv = formula.num_vars as usize;
let orig_w = original_weights(meta, orig_nv, mode);
let stages = StageReport {
simplify: Some(simplify_outcome(config)),
..StageReport::default()
};
let purpose = if weighted {
SimplifyPurpose::WeightedCount
} else {
SimplifyPurpose::Count
};
let mut simplified = simplify(formula, &preprocess_config(config, purpose, &orig_w));
let mut telemetry = PreprocessTelemetry::from_simplified(&simplified, config.stages.simplify);
if weighted {
let folded = weighted_lift::folded_weights(&simplified, &orig_w);
if let DveVerdict::Revert(why) =
weighted_lift::dve_verdict(&simplified, &folded, true)
{
diag!("c note: reverting dve ({why})");
simplified.dve_reduced = None;
}
}
if !crate::cnf::contains_empty_clause(&simplified.reduced_formula().clauses)
&& simplified.reduced_formula().num_vars == 0
{
simplified.promote_all_backbone_to_live();
}
let folded = weighted_lift::folded_weights(&simplified, &orig_w);
let stage1_weights = simplified.composed_var_map().carry_weights(&folded);
let stage1_lift = if weighted {
weighted_lift::weighted_lift(&simplified, &orig_w, &folded)
} else {
BigRational::one()
};
telemetry.total_ms = started.elapsed().as_millis() as u64;
CountStage1 {
simplified,
stage1_weights,
stage1_lift,
stages,
telemetry,
mode,
}
}
pub(super) fn finish_count_preserving_attempt(
stage1: &CountStage1,
config: &RunConfig,
) -> Result<PreprocessBundle, VitriError> {
let weighted = stage1.mode.is_weighted();
let started = std::time::Instant::now();
let mut bundle = finish_count_preserving_attempt_using(
stage1,
config,
|formula, weights, report, telemetry| {
if weighted {
weighted_arjun_stage(formula, weights, config, report, telemetry)
} else {
plain_arjun_stage(formula, config, report, telemetry)
}
},
)?;
bundle.telemetry.total_ms = stage1
.telemetry
.total_ms
.saturating_add(started.elapsed().as_millis() as u64);
Ok(bundle)
}
fn finish_count_preserving_attempt_using(
stage1: &CountStage1,
config: &RunConfig,
run_arjun: impl FnOnce(
&CnfFormula,
&Weights<Reduced>,
&mut StageReport,
&mut PreprocessTelemetry,
) -> Result<CountArjun, VitriError>,
) -> Result<PreprocessBundle, VitriError> {
let simplified = &stage1.simplified;
let mut stages = stage1.stages.clone();
let mut telemetry = stage1.telemetry;
if let Some(mut bundle) = refuted(
&simplified.reduced_formula().clauses,
simplified.original.num_vars,
stage1.mode,
None,
stages.clone(),
telemetry,
) {
bundle.decision_trace = simplified.decision_trace.clone();
return Ok(bundle);
}
let arjun_input = config
.retain_arjun_input
.then(|| simplified.reduced_formula().clone());
let arjun = run_arjun(
simplified.reduced_formula(),
&stage1.stage1_weights,
&mut stages,
&mut telemetry,
)?;
if let Some(f) = arjun.reduced_formula()
&& let Some(mut bundle) = refuted(
&f.clauses,
simplified.original.num_vars,
stage1.mode,
None,
stages.clone(),
telemetry,
)
{
bundle.decision_trace = simplified.decision_trace.clone();
return Ok(bundle);
}
let count_lift = if stage1.mode.is_weighted() {
CountLift::default()
} else {
CountLift {
simplify_pow2: simplified.count_lift_pow2(0),
arjun_pow2: arjun_multiplier_exp(&arjun),
}
};
let record = count_preserving_record(
simplified,
&arjun,
simplified.original.num_vars,
stage1.mode,
&stage1.stage1_lift,
&stage1.stage1_weights,
count_lift,
);
let (reduced, learnt_clauses_reduced_dimacs, independent_support_reduced) = match arjun {
CountArjun::Plain(ar) => (ar.formula, ar.learnt_clauses, Some(ar.independent_support)),
CountArjun::Weighted(ar) => (ar.formula, Vec::new(), None),
CountArjun::Skipped => (simplified.reduced_formula().clone(), Vec::new(), None),
};
if config.arjun.export_learned_clauses {
diag!(
"c note: exporting {} learnt clauses from arjun",
learnt_clauses_reduced_dimacs.len(),
);
}
Ok(PreprocessBundle {
reduced,
record,
learnt_clauses_reduced_dimacs,
stages,
count_lift,
telemetry,
decision_trace: simplified.decision_trace.clone(),
arjun_input,
independent_support_reduced,
})
}
fn arjun_multiplier_exp(arjun: &CountArjun) -> u32 {
match arjun {
CountArjun::Plain(ar) => ar.multiplier_exp,
CountArjun::Weighted(_) | CountArjun::Skipped => 0,
}
}
pub(super) type CountArjun = ArjunOutcome<ArjunResult, ArjunWeightedResult>;
pub(super) fn grew_clause_count(
baseline_clauses: usize,
reduced: &CnfFormula,
) -> Option<DiscardReason> {
(!arjun_keep_reduction(ArjunKeep::ClauseCount {
raw_clauses: baseline_clauses,
reduced_clauses: reduced.clauses.len(),
}))
.then_some(DiscardReason::NotSmaller)
}
pub(super) fn plain_arjun_stage(
formula: &CnfFormula,
config: &RunConfig,
report: &mut StageReport,
telemetry: &mut PreprocessTelemetry,
) -> Result<CountArjun, VitriError> {
let ar = arjun_stage(
formula,
config,
report,
telemetry,
|budget, no_sbva| run_arjun_anytime(formula, budget, config.arjun, no_sbva),
|ar| {
grew_clause_count(
config
.arjun_clause_growth
.clause_count_baseline(formula.clauses.len()),
&ar.formula,
)
},
)?;
Ok(ar.map_or(CountArjun::Skipped, CountArjun::Plain))
}
pub(super) fn weighted_arjun_stage(
formula: &CnfFormula,
weights: &Weights<Reduced>,
config: &RunConfig,
report: &mut StageReport,
telemetry: &mut PreprocessTelemetry,
) -> Result<CountArjun, VitriError> {
let ar = arjun_stage(
formula,
config,
report,
telemetry,
|budget, no_sbva| {
run_arjun_weighted_anytime(
formula,
&weights.to_dimacs_pairs(),
budget,
config.arjun,
no_sbva,
)
},
|ar| {
if !arjun_keep_reduction(ArjunKeep::weighted_for(formula.num_vars, ar)) {
return Some(DiscardReason::WeightedUnusable);
}
grew_clause_count(
config
.arjun_clause_growth
.clause_count_baseline(formula.clauses.len()),
&ar.formula,
)
},
)?;
Ok(ar.map_or(CountArjun::Skipped, CountArjun::Weighted))
}
pub(super) fn count_preserving_record(
simplified: &SimplifiedFormula,
arjun: &CountArjun,
original_num_vars: u32,
mode: Mode,
stage1_lift: &BigRational,
stage1_weights: &Weights<Reduced>,
count_lift: CountLift,
) -> PreprocessRecord {
let weighted = mode.is_weighted();
let stage1_to_original =
|j: usize| VarId(simplified.reduced_var_to_original(j) as u32).to_dimacs();
let reduced_to_original_dimacs = match arjun.var_map() {
Some(input_to_reduced) => input_to_reduced.invert_composed(
arjun.reduced_formula().unwrap().num_vars,
stage1_to_original,
),
None => simplified.composed_var_map(),
};
let (mut forced_literals_original_dimacs, mut free_vars_original_dimacs) =
simplified.stripped_forced_and_free();
if let Some(dve) = simplified.dve_reduced.as_ref() {
for j in dve.free_vars() {
free_vars_original_dimacs
.push(VarId(simplified.pre_dve_var_to_original(j) as u32).to_dimacs() as u32);
}
}
let mut arjun_rational = BigRational::one();
let mut final_weights: Option<Weights<Reduced>> = weighted.then(|| stage1_weights.clone());
match arjun {
CountArjun::Plain(ar) => {
for l in &ar.backbone {
let j = l.var.0 as usize;
if ar.input_to_reduced_lit.get(l.var).is_some() {
continue;
}
let o = stage1_to_original(j);
forced_literals_original_dimacs.push(if l.positive { o } else { -o });
}
}
CountArjun::Weighted(ar) => {
arjun_rational = ar.multiplier.clone();
final_weights = Some(ar.weights.clone());
}
CountArjun::Skipped => {}
}
let lift = if weighted {
RecordLift::Weight(stage1_lift * arjun_rational)
} else {
RecordLift::Pow2(count_lift.total_pow2())
};
PreprocessRecord {
forced_literals_original_dimacs,
free_vars_original_dimacs,
reduced_weights: final_weights.as_ref().map(Weights::to_record_rows),
..PreprocessRecord::new(mode, lift, original_num_vars, reduced_to_original_dimacs)
}
}
#[cfg(test)]
mod tests;