use std::collections::HashSet;
use crate::cyclotomic::IsRing;
use crate::enumerate::boundary::Boundary;
use crate::enumerate::canonical::CanonicalOps;
use crate::enumerate::dfs::{hashset_recorder, rat_enum_step};
use crate::enumerate::prune::Prunes;
use crate::enumerate::stats::DfsStats;
pub fn branch_factor(hturn: i8, step: i8) -> usize {
let hm1 = (hturn.max(1) - 1) as usize;
2 * (hm1 / step.max(1) as usize) + 1
}
pub fn splitting_depth(n_threads: usize, branching: usize) -> usize {
if n_threads <= 1 || branching <= 1 {
return 0;
}
let target = (10 * n_threads) as f64;
let depth = (target.ln() / (branching as f64).ln()).ceil() as usize;
depth.max(1)
}
#[allow(clippy::too_many_arguments)]
pub fn rat_enum_parallel<ZZ, B, Mk>(
mk: Mk,
max_steps: usize,
step: i8,
n_threads: usize,
ops: CanonicalOps,
label: &str,
prefix: &str,
paranoid: bool,
prunes: &Prunes,
) -> (Vec<Vec<i8>>, DfsStats)
where
ZZ: IsRing + Sync,
B: Boundary<ZZ>,
Mk: Fn(&[i8]) -> B + Sync,
{
let branching = branch_factor(ZZ::hturn(), step);
let split_depth = splitting_depth(n_threads, branching);
println!("-------- {label} started --------");
if paranoid {
println!("paranoid: per-step fresh-snake cross-check enabled");
}
println!("parallel: n_threads={n_threads} branching={branching} split_depth={split_depth}");
let mut closed_main: HashSet<Vec<i8>> = HashSet::new();
let mut seeds: Vec<Vec<i8>> = Vec::new();
let mut seed_stats = DfsStats::default();
{
let mut b = mk(&[]);
let mut record_closed = hashset_recorder(&mut closed_main);
rat_enum_step::<ZZ, B>(
&mut b,
max_steps,
step,
&mut record_closed,
&mut seed_stats,
ops,
paranoid,
prunes,
None,
split_depth,
&mut seeds,
);
}
println!("parallel: {} seed states collected", seeds.len());
let (merged, worker_stats) = parallel_drain_seeds::<ZZ, B, Mk>(
&mk,
&seeds,
closed_main,
seed_stats,
max_steps,
step,
n_threads,
ops,
paranoid,
prunes,
);
println!(
"-------- {label} completed --------\n{prefix}{} rats found",
merged.len()
);
let mut result: Vec<Vec<i8>> = merged.into_iter().collect();
result.sort_by_key(|x| x.len());
(result, worker_stats)
}
#[allow(clippy::too_many_arguments)]
pub fn parallel_drain_seeds<ZZ, B, Mk>(
mk: &Mk,
seeds: &[Vec<i8>],
closed_main: HashSet<Vec<i8>>,
seed_stats: DfsStats,
max_steps: usize,
step: i8,
n_threads: usize,
ops: CanonicalOps,
paranoid: bool,
prunes: &Prunes,
) -> (HashSet<Vec<i8>>, DfsStats)
where
ZZ: IsRing + Sync,
B: Boundary<ZZ>,
Mk: Fn(&[i8]) -> B + Sync,
{
let (workers_set, workers_stats) = crate::util::parallel::parallel_drain(
seeds.len(),
n_threads,
|| (HashSet::<Vec<i8>>::new(), DfsStats::default()),
|acc, i| {
let mut b = mk(&seeds[i]);
let mut record = hashset_recorder(&mut acc.0);
rat_enum_step::<ZZ, B>(
&mut b,
max_steps,
step,
&mut record,
&mut acc.1,
ops,
paranoid,
prunes,
None,
usize::MAX,
&mut Vec::new(),
);
},
|(mut sa, mut sta), (sb, stb)| {
sa.extend(sb);
sta.merge(&stb);
(sa, sta)
},
);
let mut merged = closed_main;
merged.extend(workers_set);
let mut total_stats = seed_stats;
total_stats.merge(&workers_stats);
(merged, total_stats)
}