use std::collections::BinaryHeap;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use dashmap::mapref::entry::Entry as MapEntry;
use dashmap::DashMap;
use rustc_hash::{FxBuildHasher, FxHashMap};
use crate::distances::Distances;
use crate::field::{Coeffs, Entry, HeapEntry};
use crate::reduce::{Engine, PairScratch, Pivots};
use crate::simplex::Simplex;
use crate::Bar;
#[derive(Clone)]
struct Owner {
coeff: u64,
diameter: f64,
col: usize,
v: Arc<[Entry]>,
}
type Table = DashMap<u64, Owner, FxBuildHasher>;
enum Claim {
Won,
Displaced(usize),
Lost,
}
fn claim(table: &Table, index: u64, owner: Owner) -> Claim {
match table.entry(index) {
MapEntry::Vacant(slot) => {
slot.insert(owner);
Claim::Won
}
MapEntry::Occupied(mut slot) => {
let incumbent = slot.get().col;
if incumbent < owner.col {
Claim::Lost
} else {
slot.insert(owner);
Claim::Displaced(incumbent)
}
}
}
}
#[repr(align(128))]
struct Padded<T>(T);
const IDLE_YIELDS: u32 = 32;
const IDLE_SLEEP_DOUBLINGS: u32 = 7;
struct WorkQueue {
next: Padded<AtomicUsize>,
pending: Padded<AtomicUsize>,
has_requeued: Padded<AtomicBool>,
requeued: Mutex<Vec<usize>>,
len: usize,
}
impl WorkQueue {
fn new(len: usize) -> Self {
Self {
next: Padded(AtomicUsize::new(0)),
pending: Padded(AtomicUsize::new(len)),
has_requeued: Padded(AtomicBool::new(false)),
requeued: Mutex::new(Vec::new()),
len,
}
}
fn take(&self) -> Option<usize> {
if self.has_requeued.0.load(Ordering::Acquire) {
if let Some(col) = self.pop_requeued() {
return Some(col);
}
}
self.next
.0
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |i| {
(i < self.len).then_some(i + 1)
})
.ok()
}
fn pop_requeued(&self) -> Option<usize> {
let mut queue = self.requeued.lock().unwrap();
let col = queue.pop();
if queue.is_empty() {
self.has_requeued.0.store(false, Ordering::Release);
}
col
}
fn requeue(&self, col: usize) {
let mut queue = self.requeued.lock().unwrap();
queue.push(col);
self.has_requeued.0.store(true, Ordering::Release);
}
}
struct Ctx<'a> {
columns: &'a [Simplex],
dim: usize,
prev_pivots: &'a Pivots,
table: &'a Table,
queue: &'a WorkQueue,
empty_v: Arc<[Entry]>,
}
#[derive(Default)]
struct Scratch {
working_cob: BinaryHeap<HeapEntry>,
working_red: BinaryHeap<HeapEntry>,
cofacet_buf: Vec<Entry>,
v_buf: Vec<Entry>,
verts: Vec<usize>,
cofacet_verts: Vec<usize>,
pairs: PairScratch,
}
enum Pass {
Owned(Option<usize>),
Essential,
Requeue,
}
impl<C: Coeffs + Sync, D: Distances + Sync> Engine<'_, C, D> {
pub(crate) fn reduce_dimension_parallel(
&self,
columns: &[Simplex],
dim: usize,
prev_pivots: &Pivots,
budget: usize,
) -> (Pivots, Vec<Bar>) {
if columns.is_empty() {
return (FxHashMap::default(), Vec::new());
}
let (table, bars) = self.converge(columns, dim, prev_pivots, budget);
(self.to_pivots(&table), bars)
}
fn converge(
&self,
columns: &[Simplex],
dim: usize,
prev_pivots: &Pivots,
budget: usize,
) -> (Table, Vec<Bar>) {
let table: Table = DashMap::with_capacity_and_hasher(columns.len(), FxBuildHasher);
let queue = WorkQueue::new(columns.len());
let ctx = Ctx {
columns,
dim,
prev_pivots,
table: &table,
queue: &queue,
empty_v: Arc::from(Vec::new()),
};
self.install(|| {
rayon::broadcast(|worker| {
if worker.index() < budget {
let mut scratch = Scratch::default();
self.worker(&ctx, &mut scratch);
}
});
});
let bars = self.collect_bars(&ctx);
drop(ctx);
(table, bars)
}
#[cfg(test)]
pub(crate) fn parallel_pivot_registry(
&self,
columns: &[Simplex],
dim: usize,
prev_pivots: &Pivots,
budget: usize,
) -> Vec<(u64, u64, usize, u64)> {
let (table, _) = self.converge(columns, dim, prev_pivots, budget);
let mut registry: Vec<(u64, u64, usize, u64)> = table
.iter()
.map(|r| {
let owner = r.value();
(*r.key(), owner.coeff, owner.col, owner.diameter.to_bits())
})
.collect();
registry.sort_unstable();
registry
}
#[inline(never)]
fn worker(&self, ctx: &Ctx, scratch: &mut Scratch) {
let mut done = 0usize;
let mut idle = 0u32;
loop {
let Some(col) = ctx.queue.take() else {
if done > 0 {
ctx.queue.pending.0.fetch_sub(done, Ordering::AcqRel);
done = 0;
}
if ctx.queue.pending.0.load(Ordering::Acquire) == 0 {
return;
}
if idle < IDLE_YIELDS {
std::thread::yield_now();
} else {
let step = (idle - IDLE_YIELDS).min(IDLE_SLEEP_DOUBLINGS);
std::thread::sleep(std::time::Duration::from_micros(1 << step));
}
idle += 1;
continue;
};
idle = 0;
match self.reduce_pass(ctx, scratch, col) {
Pass::Owned(displaced) => {
if let Some(k) = displaced {
ctx.queue.pending.0.fetch_add(1, Ordering::AcqRel);
ctx.queue.requeue(k);
}
done += 1;
}
Pass::Essential => done += 1,
Pass::Requeue => ctx.queue.requeue(col),
}
}
}
fn reduce_pass(&self, ctx: &Ctx, scratch: &mut Scratch, col: usize) -> Pass {
let column = ctx.columns[col];
scratch.working_cob.clear();
scratch.working_red.clear();
let mut pivot = self.init_coboundary(
column,
ctx.dim,
|index| ctx.table.contains_key(&index),
&mut scratch.working_cob,
&mut scratch.cofacet_buf,
&mut scratch.verts,
&mut scratch.cofacet_verts,
&mut scratch.pairs,
);
let mut built = !(pivot.is_some() && scratch.working_cob.is_empty());
loop {
let Some(p) = pivot else {
return Pass::Essential;
};
let index = self.ops.index(p);
let held = ctx.table.get(&index).map(|r| r.clone());
match held {
Some(owner) if owner.col < col => {
if !built {
self.build_full_coboundary(
column,
ctx.dim,
&mut scratch.working_cob,
&mut scratch.verts,
);
built = true;
}
self.fold_reducer(
p,
owner.coeff,
ctx.columns[owner.col],
&owner.v,
ctx.dim,
&mut scratch.working_red,
&mut scratch.working_cob,
&mut scratch.verts,
);
pivot = self.get_pivot(&mut scratch.working_cob);
}
_ => {
if let Some(next) = self.reduce_apparent_facet(
p,
ctx.dim,
&mut scratch.working_red,
&mut scratch.working_cob,
&mut scratch.verts,
&mut scratch.pairs,
) {
pivot = next;
} else {
return self.claim_pivot(ctx, scratch, p, index, col);
}
}
}
}
}
fn claim_pivot(
&self,
ctx: &Ctx,
scratch: &mut Scratch,
p: Entry,
index: u64,
col: usize,
) -> Pass {
scratch.v_buf.clear();
self.drain_into(&mut scratch.working_red, &mut scratch.v_buf);
let v = if scratch.v_buf.is_empty() {
ctx.empty_v.clone()
} else {
Arc::from(scratch.v_buf.as_slice())
};
let owner = Owner {
coeff: self.ops.coeff(p),
diameter: p.diameter,
col,
v,
};
match claim(ctx.table, index, owner) {
Claim::Won => Pass::Owned(None),
Claim::Displaced(k) => Pass::Owned(Some(k)),
Claim::Lost => Pass::Requeue,
}
}
fn to_pivots(&self, table: &Table) -> Pivots {
table
.iter()
.map(|r| {
let owner = r.value();
(*r.key(), (owner.coeff, owner.col))
})
.collect()
}
fn collect_bars(&self, ctx: &Ctx) -> Vec<Bar> {
let mut bars = Vec::new();
let mut owns_pivot = vec![false; ctx.columns.len()];
for r in ctx.table.iter() {
let owner = r.value();
owns_pivot[owner.col] = true;
let birth = ctx.columns[owner.col].diameter;
if owner.diameter > birth {
bars.push(Bar {
dim: ctx.dim,
birth,
death: owner.diameter,
});
}
}
for (col, &owned) in owns_pivot.iter().enumerate() {
let column = ctx.columns[col];
let prior_death =
!self.params.use_clearing && ctx.prev_pivots.contains_key(&column.index);
if !owned && !prior_death {
bars.push(Bar {
dim: ctx.dim,
birth: column.diameter,
death: f64::INFINITY,
});
}
}
bars
}
}
#[cfg(test)]
mod tests {
use rustc_hash::FxHashMap;
use crate::field::Z2;
use crate::reduce::{Engine, Pivots};
use crate::simplex::Simplex;
use crate::{Diagram, DistanceMatrix, RipsParams};
fn edge_columns(
dist: &DistanceMatrix,
engine: &Engine<'_, Z2, DistanceMatrix>,
) -> Vec<Simplex> {
let mut columns = Vec::new();
for i in 1..dist.len() {
for j in 0..i {
let diameter = dist.get(i, j);
if engine.in_complex(diameter) {
columns.push(Simplex {
diameter,
index: engine.bt.get(i, 2) + j as u64,
});
}
}
}
columns.sort_unstable_by(|a, b| {
b.diameter
.total_cmp(&a.diameter)
.then(a.index.cmp(&b.index))
});
columns
}
fn points(seed: u64, n: usize, coord_dim: usize) -> Vec<Vec<f64>> {
let mut x = seed | 1;
let mut next = || {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
(x >> 11) as f64 / (1u64 << 53) as f64
};
(0..n)
.map(|_| (0..coord_dim).map(|_| next()).collect())
.collect()
}
fn grid(side: usize) -> Vec<Vec<f64>> {
(0..side)
.flat_map(|a| (0..side).map(move |b| vec![a as f64, b as f64]))
.collect()
}
fn assert_registry_is_worker_invariant(dist: &DistanceMatrix, label: &str) {
let mut serial_params = RipsParams::new(1);
serial_params.threads = 1;
let serial = Engine::new(dist, &serial_params, Z2).unwrap();
let columns = edge_columns(dist, &serial);
assert!(columns.len() > 64, "{label}: too few columns to be a gate");
let empty: Pivots = FxHashMap::default();
let mut diagram = Diagram::default();
let want = serial.reduce_dimension(&columns, 1, &empty, &mut diagram);
let mut first: Option<Vec<(u64, u64, usize, u64)>> = None;
for budget in [2usize, 3, 4, 8] {
let mut params = RipsParams::new(1);
params.threads = budget;
let engine = Engine::new(dist, ¶ms, Z2).unwrap();
let got = engine.parallel_pivot_registry(&columns, 1, &empty, budget);
let by_index: Pivots = got
.iter()
.map(|&(index, coeff, col, _)| (index, (coeff, col)))
.collect();
assert_eq!(
by_index, want,
"{label}: {budget} workers disagree with the serial pivot registry"
);
match &first {
None => first = Some(got),
Some(want) => assert_eq!(
&got, want,
"{label}: {budget} workers disagree with 2 workers, diameter bits included"
),
}
}
}
#[test]
fn pivot_registry_is_worker_invariant_on_a_cloud() {
let dist = DistanceMatrix::from_points(&points(20260818, 100, 3)).unwrap();
assert_registry_is_worker_invariant(&dist, "cloud(n=100,d=3)");
}
#[test]
fn pivot_registry_is_worker_invariant_on_ties() {
let dist = DistanceMatrix::from_points(&grid(9)).unwrap();
assert_registry_is_worker_invariant(&dist, "grid(9x9)");
}
}