use std::time::{Duration, Instant};
use rayon::prelude::*;
use super::{
build_pool, finish, mark_dirty, prepare, test_edge, tombstone, CollapseStats, CollapseTimings,
CollapsedRips, Execution, Prepared, RemovalStep, Scratch,
};
use crate::distances::Distances;
use crate::{DistanceMatrix, Result, SparseDistanceMatrix};
const WINDOW_PER_WORKER: usize = 32;
const WINDOW_CAP: usize = 4096;
type Witnesses = Vec<(f64, usize)>;
fn window_for(workers: usize) -> usize {
(WINDOW_PER_WORKER * workers.max(1)).clamp(1, WINDOW_CAP)
}
fn nanos(d: Duration) -> u64 {
u64::try_from(d.as_nanos()).unwrap_or(u64::MAX)
}
pub fn collapse_dense_ordered_parallel(
dist: &DistanceMatrix,
threshold: Option<f64>,
threads: usize,
) -> Result<CollapsedRips> {
collapse_ordered_owned(dist, threshold, threads, None)
}
pub fn collapse_sparse_ordered_parallel(
dist: &SparseDistanceMatrix,
threshold: Option<f64>,
threads: usize,
) -> Result<CollapsedRips> {
collapse_ordered_owned(dist, threshold, threads, None)
}
#[doc(hidden)]
pub fn collapse_dense_ordered_with_window(
dist: &DistanceMatrix,
threshold: Option<f64>,
threads: usize,
window: usize,
) -> Result<CollapsedRips> {
collapse_ordered_owned(dist, threshold, threads, Some(window.max(1)))
}
#[doc(hidden)]
pub fn collapse_sparse_ordered_with_window(
dist: &SparseDistanceMatrix,
threshold: Option<f64>,
threads: usize,
window: usize,
) -> Result<CollapsedRips> {
collapse_ordered_owned(dist, threshold, threads, Some(window.max(1)))
}
fn collapse_ordered_owned<D: Distances + Sync>(
dist: &D,
threshold: Option<f64>,
threads: usize,
window: Option<usize>,
) -> Result<CollapsedRips> {
if threads.max(1) == 1 {
return serial(dist, threshold);
}
let pool = build_pool(threads)?;
let w = window.unwrap_or_else(|| window_for(pool.current_num_threads()));
collapse_ordered_core(dist, threshold, Some(&pool), w)
}
pub(crate) fn collapse_ordered_in<D: Distances + Sync>(
dist: &D,
threshold: Option<f64>,
pool: Option<&rayon::ThreadPool>,
) -> Result<CollapsedRips> {
match pool.filter(|p| p.current_num_threads() > 1) {
Some(p) => collapse_ordered_core(
dist,
threshold,
Some(p),
window_for(p.current_num_threads()),
),
None => serial(dist, threshold),
}
}
fn serial<D: Distances>(dist: &D, threshold: Option<f64>) -> Result<CollapsedRips> {
super::collapse_impl(dist, threshold)
}
fn collapse_ordered_core<D: Distances + Sync>(
dist: &D,
threshold: Option<f64>,
pool: Option<&rayon::ThreadPool>,
window: usize,
) -> Result<CollapsedRips> {
let Prepared {
mut edges,
mut adj,
run,
} = prepare(dist, threshold)?;
let window = window.max(1);
let mut stats = CollapseStats::new(edges.len());
let mut steps: Vec<RemovalStep> = Vec::new();
let mut scratch = Scratch::default();
let mut dirty: Vec<bool> = vec![false; edges.len()];
let mut test_all = true;
let mut members: Vec<usize> = Vec::new();
let mut stale: Vec<bool> = Vec::new();
let mut predicate_time = Duration::ZERO;
let mut retirement_time = Duration::ZERO;
let mut repair_time = Duration::ZERO;
loop {
stats.epochs += 1;
let mut removed_any = false;
let mut test_all_next = false;
let mut c = 0usize;
while c < edges.len() {
members.clear();
let mut scan = c;
while scan < edges.len() && members.len() < window {
if edges[scan].alive && (test_all || dirty[scan]) {
members.push(scan);
}
scan += 1;
}
let Some(&last) = members.last() else {
break;
};
stats.window_batches += 1;
stats.window_slots_offered = stats.window_slots_offered.saturating_add(window);
stats.window_members_formed += members.len();
stats.edge_tests += members.len();
let predicate_start = Instant::now();
let mut cached: Vec<(Option<Witnesses>, usize)> = match pool {
Some(pool) => pool.install(|| {
members
.par_iter()
.map_init(Scratch::default, |s, &idx| {
let e = &edges[idx];
let w = test_edge(&adj, e.u, e.v, e.value, run.terminal, s);
(w, s.cands.len())
})
.collect()
}),
None => members
.iter()
.map(|&idx| {
let e = &edges[idx];
let w = test_edge(&adj, e.u, e.v, e.value, run.terminal, &mut scratch);
(w, scratch.cands.len())
})
.collect(),
};
predicate_time += predicate_start.elapsed();
for &(_, k) in &cached {
stats.max_common_neighborhood = stats.max_common_neighborhood.max(k);
}
stale.clear();
stale.resize(members.len(), false);
let retirement_start = Instant::now();
let mut next = 0usize;
for j in c..=last {
let slot = if next < members.len() && members[next] == j {
next += 1;
Some(next - 1)
} else {
None
};
if !edges[j].alive || !(test_all || dirty[j]) {
continue;
}
stats.logical_tests += 1;
dirty[j] = false;
let (u, v, value) = (edges[j].u, edges[j].v, edges[j].value);
let witnesses = match slot.filter(|&k| !stale[k]) {
Some(k) => {
stats.window_members_reused += 1;
cached[k].0.take()
}
None => {
stats.edge_tests += 1;
let repairing = slot.is_some();
let repair_start = repairing.then(Instant::now);
let w = test_edge(&adj, u, v, value, run.terminal, &mut scratch);
if let Some(start) = repair_start {
repair_time += start.elapsed();
}
stats.max_common_neighborhood =
stats.max_common_neighborhood.max(scratch.cands.len());
w
}
};
let Some(witnesses) = witnesses else {
continue;
};
edges[j].alive = false;
if mark_dirty(&adj, &mut dirty, u, v, &mut scratch) {
let set = &scratch.marks;
for (k, &pos) in members.iter().enumerate().skip(next) {
if stale[k] {
continue;
}
let m = &edges[pos];
if set.binary_search(&m.u).is_ok() && set.binary_search(&m.v).is_ok() {
stale[k] = true;
stats.invalidated_results += 1;
}
}
} else {
test_all_next = true;
let mut dropped = 0usize;
for st in stale.iter_mut().skip(next) {
if !*st {
*st = true;
dropped += 1;
}
}
if dropped > 0 {
stats.global_invalidations += 1;
stats.invalidated_results += dropped;
}
}
tombstone(&mut adj, u, v);
stats.witness_segments += witnesses.len();
steps.push(RemovalStep {
u,
v,
value,
epoch: stats.epochs,
witnesses,
});
removed_any = true;
}
retirement_time += retirement_start.elapsed();
c = last + 1;
}
if !removed_any {
break;
}
test_all = test_all_next;
}
let timings = CollapseTimings {
predicate_ns: nanos(predicate_time),
retirement_ns: nanos(retirement_time),
repair_ns: nanos(repair_time),
};
finish(run, Execution::Ordered, &edges, steps, stats, timings)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::collapse::{collapse_dense, collapse_sparse, CollapseCertificate};
const WINDOWS: [usize; 5] = [1, 2, 8, 64, 10_000];
fn edge_bits(m: &SparseDistanceMatrix) -> Vec<(usize, usize, u64)> {
m.edges().map(|(u, v, d)| (u, v, d.to_bits())).collect()
}
fn assert_certificate_bits(a: &CollapseCertificate, b: &CollapseCertificate, label: &str) {
assert_eq!(a.algorithm_version(), b.algorithm_version(), "{label}");
assert_eq!(a.vertex_count(), b.vertex_count(), "{label}");
assert_eq!(
a.requested_threshold().map(f64::to_bits),
b.requested_threshold().map(f64::to_bits),
"{label}"
);
assert_eq!(
a.terminal_level().to_bits(),
b.terminal_level().to_bits(),
"{label}"
);
assert_eq!(a.input_edge_count(), b.input_edge_count(), "{label}");
assert_eq!(a.output_edge_count(), b.output_edge_count(), "{label}");
assert_eq!(a.steps().len(), b.steps().len(), "{label}");
for (x, y) in a.steps().iter().zip(b.steps()) {
assert_eq!(x.edge(), y.edge(), "{label}");
assert_eq!(x.value().to_bits(), y.value().to_bits(), "{label}");
assert_eq!(x.epoch(), y.epoch(), "{label}");
let wx: Vec<_> = x
.witnesses()
.iter()
.map(|&(t, w)| (t.to_bits(), w))
.collect();
let wy: Vec<_> = y
.witnesses()
.iter()
.map(|&(t, w)| (t.to_bits(), w))
.collect();
assert_eq!(wx, wy, "{label}");
}
}
fn assert_occupancy_bounds(r: &CollapsedRips, label: &str) {
let s = &r.stats;
assert!(
s.window_members_formed <= s.window_slots_offered,
"{label}: formed {} > offered {}",
s.window_members_formed,
s.window_slots_offered
);
assert!(
s.window_members_reused <= s.window_members_formed,
"{label}: reused {} > formed {}",
s.window_members_reused,
s.window_members_formed
);
assert_eq!(
s.window_members_reused + s.invalidated_results,
s.window_members_formed,
"{label}: reused {} plus repairs {} != formed {}",
s.window_members_reused,
s.invalidated_results,
s.window_members_formed
);
}
fn assert_matches_serial(ordered: &CollapsedRips, serial: &CollapsedRips, label: &str) {
assert_occupancy_bounds(ordered, label);
assert_eq!(ordered.certificate, serial.certificate, "{label}");
assert_certificate_bits(&ordered.certificate, &serial.certificate, label);
assert_eq!(
edge_bits(&ordered.matrix),
edge_bits(&serial.matrix),
"{label}"
);
assert_eq!(ordered.stats.epochs, serial.stats.epochs, "{label}");
assert_eq!(
ordered.stats.logical_tests, serial.stats.edge_tests,
"{label}"
);
assert_eq!(
ordered.stats.input_edges, serial.stats.input_edges,
"{label}"
);
assert_eq!(
ordered.stats.output_edges, serial.stats.output_edges,
"{label}"
);
assert_eq!(
ordered.stats.removed_edges, serial.stats.removed_edges,
"{label}"
);
assert_eq!(
ordered.stats.witness_segments, serial.stats.witness_segments,
"{label}"
);
assert!(
ordered.stats.edge_tests >= ordered.stats.logical_tests,
"{label}"
);
}
fn tie_heavy_20() -> DistanceMatrix {
let mut condensed = Vec::new();
for i in 1..20usize {
for j in 0..i {
condensed.push(((i * j + i + j) % 5 + 1) as f64);
}
}
DistanceMatrix::from_condensed(condensed).unwrap()
}
fn unit_k4() -> DistanceMatrix {
DistanceMatrix::from_condensed(vec![1.0; 6]).unwrap()
}
fn random_two_value_40() -> DistanceMatrix {
let mut state = 0x9e37_79b9_7f4a_7c15u64;
let mut condensed = Vec::new();
for _ in 0..40 * 39 / 2 {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
condensed.push(if (state >> 33) % 3 == 0 { 1.0 } else { 2.0 });
}
DistanceMatrix::from_condensed(condensed).unwrap()
}
fn four_cycle() -> SparseDistanceMatrix {
SparseDistanceMatrix::from_triplets(
4,
&[(0, 1, 1.0), (1, 2, 1.0), (2, 3, 1.0), (0, 3, 1.0)],
)
.unwrap()
}
fn two_unit_k4s() -> SparseDistanceMatrix {
let mut triplets = Vec::new();
for base in [0usize, 4] {
for v in 1..4 {
for u in 0..v {
triplets.push((base + u, base + v, 1.0));
}
}
}
SparseDistanceMatrix::from_triplets(8, &triplets).unwrap()
}
#[test]
fn dense_ordered_matches_serial_across_threads_and_windows() {
for (name, d) in [
("tie_heavy_20", tie_heavy_20()),
("k4", unit_k4()),
("random_40", random_two_value_40()),
] {
let base = collapse_dense(&d, None).unwrap();
assert!(base.stats.removed_edges > 0, "{name}");
for w in WINDOWS {
let r = collapse_ordered_core(&d, None, None, w).unwrap();
assert_matches_serial(&r, &base, &format!("{name} pool=none w={w}"));
assert!(r.stats.window_batches > 0);
for t in [2usize, 4] {
let r = collapse_dense_ordered_with_window(&d, None, t, w).unwrap();
assert_matches_serial(&r, &base, &format!("{name} threads={t} w={w}"));
}
}
}
}
#[test]
fn sparse_ordered_matches_serial_across_threads_and_windows() {
for (name, m) in [("four_cycle", four_cycle()), ("two_k4s", two_unit_k4s())] {
let base = collapse_sparse(&m, None).unwrap();
for w in WINDOWS {
let r = collapse_ordered_core(&m, None, None, w).unwrap();
assert_matches_serial(&r, &base, &format!("{name} pool=none w={w}"));
for t in [2usize, 4] {
let r = collapse_sparse_ordered_with_window(&m, None, t, w).unwrap();
assert_matches_serial(&r, &base, &format!("{name} threads={t} w={w}"));
}
}
}
}
#[test]
fn stale_member_repair_flips_a_verdict() {
let d = unit_k4();
let base = collapse_dense(&d, None).unwrap();
let r = collapse_dense_ordered_with_window(&d, None, 2, 64).unwrap();
assert_matches_serial(&r, &base, "k4 stale repair");
let removed: Vec<_> = r.certificate.steps().iter().map(|s| s.edge()).collect();
assert_eq!(removed, vec![(0, 1), (0, 2), (1, 2)]);
assert!(r.stats.invalidated_results >= 1);
assert!(r.stats.invalidated_results >= 1);
assert!(r.stats.edge_tests > r.stats.logical_tests);
}
#[test]
fn unit_k4_occupancy_is_exact() {
let r = collapse_dense_ordered_with_window(&unit_k4(), None, 2, 64).unwrap();
assert_eq!(r.stats.epochs, 2);
assert_eq!(r.stats.window_batches, 1);
assert_eq!(r.stats.window_slots_offered, 64);
assert_eq!(r.stats.window_members_formed, 6);
assert_eq!(r.stats.window_members_reused, 1);
assert_eq!(r.stats.invalidated_results, 5);
assert_eq!(r.stats.invalidated_results, 5);
assert_eq!(r.stats.logical_tests, 6);
assert_eq!(r.stats.edge_tests, 11);
}
#[test]
fn offered_slots_are_stages_times_the_window() {
for w in WINDOWS {
for t in [2usize, 4] {
for (name, r) in [
(
"tie_heavy_20",
collapse_dense_ordered_with_window(&tie_heavy_20(), None, t, w).unwrap(),
),
(
"random_40",
collapse_dense_ordered_with_window(&random_two_value_40(), None, t, w)
.unwrap(),
),
(
"two_k4s",
collapse_sparse_ordered_with_window(&two_unit_k4s(), None, t, w).unwrap(),
),
] {
let label = format!("{name} threads={t} w={w}");
assert!(r.stats.window_batches > 0, "{label}");
assert_eq!(
r.stats.window_slots_offered,
r.stats.window_batches * w,
"{label}"
);
assert_occupancy_bounds(&r, &label);
assert!(
r.stats.window_members_formed >= r.stats.window_batches,
"{label}"
);
}
}
}
}
#[test]
fn ordered_timings_are_measured() {
let d = random_two_value_40();
let r = collapse_dense_ordered_with_window(&d, None, 2, 8).unwrap();
assert!(r.stats.window_batches > 1);
assert!(r.timings.repair_ns <= r.timings.retirement_ns);
let k4 = collapse_dense_ordered_with_window(&unit_k4(), None, 2, 64).unwrap();
assert!(k4.stats.invalidated_results > 0);
assert!(k4.timings.repair_ns <= k4.timings.retirement_ns);
let single = collapse_dense_ordered_with_window(&d, None, 2, 1).unwrap();
assert_eq!(single.stats.invalidated_results, 0);
assert_eq!(single.timings.repair_ns, 0);
}
#[test]
fn one_worker_delegates_to_the_serial_run() {
let d = tie_heavy_20();
let base = collapse_dense(&d, None).unwrap();
for threads in [0usize, 1] {
let r = collapse_dense_ordered_parallel(&d, None, threads).unwrap();
assert_matches_serial(&r, &base, "delegation");
assert_eq!(r.stats.edge_tests, base.stats.edge_tests);
assert_eq!(r.stats.edge_tests, r.stats.logical_tests);
assert_eq!(
r.stats.max_common_neighborhood,
base.stats.max_common_neighborhood
);
assert_eq!(r.stats.invalidated_results, 0);
assert_eq!(r.stats.invalidated_results, 0);
assert_eq!(r.stats.global_invalidations, 0);
assert_eq!(r.stats.window_batches, 0);
assert_eq!(r.stats.window_slots_offered, 0);
assert_eq!(r.stats.window_members_formed, 0);
assert_eq!(r.stats.window_members_reused, 0);
assert_eq!(r.timings, CollapseTimings::default());
}
let m = two_unit_k4s();
let sparse_base = collapse_sparse(&m, None).unwrap();
let r = collapse_sparse_ordered_parallel(&m, None, 1).unwrap();
assert_matches_serial(&r, &sparse_base, "sparse delegation");
}
#[test]
fn marking_bail_invalidates_the_window_remainder() {
let d = DistanceMatrix::from_condensed(vec![1.0; 66 * 65 / 2]).unwrap();
let base = collapse_dense(&d, None).unwrap();
let r = collapse_dense_ordered_with_window(&d, None, 2, 64).unwrap();
assert_matches_serial(&r, &base, "mark bail");
assert!(r.stats.global_invalidations >= 1);
assert!(r.stats.invalidated_results >= 1);
}
#[test]
fn empty_and_tiny_inputs() {
let d0 = DistanceMatrix::from_points(&[]).unwrap();
let d1 = DistanceMatrix::from_condensed(vec![]).unwrap();
let s1 = SparseDistanceMatrix::from_triplets(1, &[]).unwrap();
for threads in [0usize, 1, 4] {
for r in [
collapse_dense_ordered_parallel(&d0, None, threads).unwrap(),
collapse_dense_ordered_parallel(&d1, None, threads).unwrap(),
collapse_sparse_ordered_parallel(&s1, None, threads).unwrap(),
] {
assert_eq!(r.certificate.algorithm_version(), 1);
assert_eq!(r.certificate.input_edge_count(), 0);
assert_eq!(r.certificate.terminal_level(), 0.0);
assert!(r.certificate.steps().is_empty());
assert_eq!(r.stats.epochs, 1);
assert_eq!(r.stats.edge_tests, 0);
assert_eq!(r.stats.logical_tests, 0);
assert_eq!(r.stats.window_batches, 0);
assert_eq!(r.stats.window_slots_offered, 0);
assert_eq!(r.stats.window_members_formed, 0);
assert_eq!(r.stats.window_members_reused, 0);
assert_eq!(r.timings, CollapseTimings::default());
}
}
}
#[test]
fn invalid_thresholds_are_rejected() {
let d = DistanceMatrix::from_condensed(vec![1.0]).unwrap();
let m = four_cycle();
for threads in [1usize, 4] {
assert!(collapse_dense_ordered_parallel(&d, Some(-1.0), threads).is_err());
assert!(collapse_dense_ordered_parallel(&d, Some(f64::NAN), threads).is_err());
assert!(collapse_sparse_ordered_parallel(&m, Some(-1.0), threads).is_err());
assert!(collapse_sparse_ordered_parallel(&m, Some(f64::NAN), threads).is_err());
assert!(collapse_dense_ordered_with_window(&d, Some(-1.0), threads, 4).is_err());
assert!(collapse_sparse_ordered_with_window(&m, Some(f64::NAN), threads, 4).is_err());
}
}
}