1use std::collections::{BTreeMap, BTreeSet};
2
3use crate::classes::{basis_class_id, canonical_space_basis, group_id, validate_h1_cocycle};
4use crate::factorization::{
5 FactorizationSummary, ProgramBlock, ProgramDecompositionSummary, program_blocks,
6};
7use crate::{
8 AtlasArtifact, Bar, CertificateLimits, Cocycle, CocycleTerm, CriticalPair, EdgeKey, Error,
9 ExplainedDiagram, PersistentClass, PersistentClassSpace, Result, RipsParams,
10 SparseDistanceMatrix, rips_persistence_sparse,
11};
12
13use super::model::{PersistenceProgram, ProgramAtomInfo, ProgramAtomState, ProgramSummary};
14use super::topology::{
15 bar_bits_equal, critical_pair_key, critical_pair_order, diagram_bits_equal, h0_diagram,
16 h0_provenance, map_critical_pair, previous_float, terminal_level,
17};
18
19impl PersistenceProgram {
20 pub fn compile(
25 input: &SparseDistanceMatrix,
26 params: &RipsParams,
27 limits: CertificateLimits,
28 ) -> Result<Self> {
29 if params.max_dim != 1 {
30 return Err(Error::InvalidInput(
31 "a persistence program requires max_dim equal to 1".into(),
32 ));
33 }
34 let (factorization, blocks) = program_blocks(input, params.threshold)?;
35 let (atoms, states) = compile_atoms(input, params, limits, &blocks)?;
36 let result = compose_result(input, params, &states)?;
37 let expected = rips_persistence_sparse(input, params)?;
38 if !diagram_bits_equal(&result.diagram, &expected) {
39 return Err(Error::InvalidInput(format!(
40 "compositional and monolithic diagrams differ: expected {:?}, got {:?}",
41 expected.bars, result.diagram.bars
42 )));
43 }
44 let topology: Vec<_> = input.edges().map(|(u, v, _)| EdgeKey::new(u, v)).collect();
45 let threshold = params.threshold.unwrap_or(f64::INFINITY);
46 let active = input
47 .edges()
48 .map(|(_, _, value)| value <= threshold)
49 .collect();
50 let summary = program_summary(factorization, &atoms, &states);
51 let separator_edges = wider_separator_edges(&atoms);
52 let (h0_deaths, h0_essential, _) = h0_provenance(input, params.threshold);
53 Ok(Self {
54 params: params.clone(),
55 limits,
56 graph: input.clone(),
57 topology,
58 active,
59 separator_edges,
60 h0_deaths,
61 h0_essential,
62 atoms,
63 states,
64 summary,
65 result,
66 })
67 }
68
69 pub(crate) fn from_verified_parts(
70 input: &SparseDistanceMatrix,
71 params: &RipsParams,
72 limits: CertificateLimits,
73 atoms: Vec<ProgramAtomInfo>,
74 states: Vec<ProgramAtomState>,
75 result: ExplainedDiagram,
76 ) -> Self {
77 let topology: Vec<_> = input.edges().map(|(u, v, _)| EdgeKey::new(u, v)).collect();
78 let threshold = params.threshold.unwrap_or(f64::INFINITY);
79 let active = input
80 .edges()
81 .map(|(_, _, value)| value <= threshold)
82 .collect();
83 let factorization = FactorizationSummary {
84 blocks: atoms.len(),
85 cyclic_blocks: atoms.iter().filter(|atom| atom.cyclic).count(),
86 bridge_edges: atoms.iter().filter(|atom| !atom.cyclic).count(),
87 cyclic_edges: atoms
88 .iter()
89 .filter(|atom| atom.cyclic)
90 .map(|atom| atom.edges.len())
91 .sum(),
92 largest_cyclic_block_edges: atoms
93 .iter()
94 .filter(|atom| atom.cyclic)
95 .map(|atom| atom.edges.len())
96 .max()
97 .unwrap_or(0),
98 };
99 let (articulation_vertices, zero_simplex_separators, widest_separator) =
100 separator_stats(&atoms);
101 let summary = program_summary(
102 ProgramDecompositionSummary {
103 articulation: factorization,
104 articulation_vertices,
105 zero_simplex_separators,
106 widest_separator,
107 separator_candidates_checked: 0,
108 separator_search_complete: true,
109 },
110 &atoms,
111 &states,
112 );
113 let separator_edges = wider_separator_edges(&atoms);
114 let (h0_deaths, h0_essential, _) = h0_provenance(input, params.threshold);
115 Self {
116 params: params.clone(),
117 limits,
118 graph: input.clone(),
119 topology,
120 active,
121 separator_edges,
122 h0_deaths,
123 h0_essential,
124 atoms,
125 states,
126 summary,
127 result,
128 }
129 }
130}
131
132fn compile_atoms(
133 input: &SparseDistanceMatrix,
134 params: &RipsParams,
135 limits: CertificateLimits,
136 blocks: &[ProgramBlock],
137) -> Result<(Vec<ProgramAtomInfo>, Vec<ProgramAtomState>)> {
138 let atoms = atom_infos(input, blocks);
139 let topology: Vec<_> = input.edges().map(|(u, v, _)| EdgeKey::new(u, v)).collect();
140 let mut states = Vec::new();
141 for atom in &atoms {
142 if !atom.cyclic {
143 continue;
144 }
145 let local = local_matrix(&atom.vertices, &atom.edges, input)?;
146 let (artifact, _) = AtlasArtifact::compile(&local, &atom_params(params), limits)
147 .map_err(|error| Error::InvalidInput(error.to_string()))?;
148 let region = artifact
149 .reduction_certificate()
150 .compile_region(&local, limits)?;
151 states.push(ProgramAtomState {
152 info_index: atom.id,
153 vertices: atom.vertices.clone(),
154 edges: atom.edges.clone(),
155 edge_positions: atom
156 .edges
157 .iter()
158 .map(|edge| {
159 topology
160 .binary_search(edge)
161 .expect("atom edge is in the program topology")
162 })
163 .collect(),
164 explained: artifact.explained().clone(),
165 artifact,
166 certified_graph: local,
167 region,
168 });
169 }
170 Ok((atoms, states))
171}
172
173pub(crate) fn atom_infos(
174 input: &SparseDistanceMatrix,
175 blocks: &[ProgramBlock],
176) -> Vec<ProgramAtomInfo> {
177 let mut counts = vec![0usize; input.len()];
178 for block in blocks {
179 for &vertex in &block.vertices {
180 counts[vertex] += 1;
181 }
182 }
183 let atoms: Vec<_> = blocks
184 .iter()
185 .enumerate()
186 .map(|(id, block)| ProgramAtomInfo {
187 id,
188 vertices: block.vertices.clone(),
189 edges: block
190 .edges
191 .iter()
192 .map(|&[u, v]| EdgeKey::new(u, v))
193 .collect(),
194 separator_vertices: block
195 .vertices
196 .iter()
197 .copied()
198 .filter(|&vertex| counts[vertex] > 1)
199 .collect(),
200 cyclic: block.edges.len() >= block.vertices.len(),
201 })
202 .collect();
203 atoms
204}
205
206pub(crate) fn atom_params(params: &RipsParams) -> RipsParams {
207 let mut atom = RipsParams::new(1).with_modulus(params.modulus);
208 atom.threshold = params.threshold;
209 atom
210}
211
212pub(crate) fn local_matrix(
213 vertices: &[usize],
214 edges: &[EdgeKey],
215 input: &SparseDistanceMatrix,
216) -> Result<SparseDistanceMatrix> {
217 let triplets: Vec<_> = edges
218 .iter()
219 .map(|edge| {
220 let u = vertices
221 .binary_search(&edge.u)
222 .expect("atom contains its edge endpoint");
223 let v = vertices
224 .binary_search(&edge.v)
225 .expect("atom contains its edge endpoint");
226 (u, v, input.get(edge.u, edge.v))
227 })
228 .collect();
229 SparseDistanceMatrix::from_triplets(vertices.len(), &triplets)
230}
231
232fn program_summary(
233 decomposition: ProgramDecompositionSummary,
234 atoms: &[ProgramAtomInfo],
235 states: &[ProgramAtomState],
236) -> ProgramSummary {
237 debug_assert!(decomposition.articulation.blocks <= atoms.len());
238 ProgramSummary {
239 atoms: atoms.len(),
240 cyclic_atoms: atoms.iter().filter(|atom| atom.cyclic).count(),
241 articulation_vertices: decomposition.articulation_vertices,
242 zero_simplex_separators: decomposition.zero_simplex_separators,
243 widest_separator: decomposition.widest_separator,
244 separator_candidates_checked: decomposition.separator_candidates_checked,
245 separator_search_complete: decomposition.separator_search_complete,
246 largest_cyclic_atom_edges: atoms
247 .iter()
248 .filter(|atom| atom.cyclic)
249 .map(|atom| atom.edges.len())
250 .max()
251 .unwrap_or(0),
252 complete_guards: states
253 .iter()
254 .map(|state| state.region.complete_guards().len())
255 .sum(),
256 guards: states.iter().map(|state| state.region.guards().len()).sum(),
257 }
258}
259
260fn separator_stats(atoms: &[ProgramAtomInfo]) -> (usize, usize, usize) {
261 let mut separators = BTreeSet::new();
262 let mut articulations = BTreeSet::new();
263 for (position, left) in atoms.iter().enumerate() {
264 for right in &atoms[position + 1..] {
265 let intersection: Vec<_> = left
266 .vertices
267 .iter()
268 .copied()
269 .filter(|vertex| right.vertices.binary_search(vertex).is_ok())
270 .collect();
271 if intersection.len() > 1 {
272 separators.insert(intersection);
273 } else if let Some(&vertex) = intersection.first() {
274 articulations.insert(vertex);
275 }
276 }
277 }
278 let widest = separators.iter().map(Vec::len).max().unwrap_or(1);
279 (articulations.len(), separators.len(), widest)
280}
281
282fn wider_separator_edges(atoms: &[ProgramAtomInfo]) -> Vec<EdgeKey> {
283 let mut edges = BTreeSet::new();
284 for (position, left) in atoms.iter().enumerate() {
285 for right in &atoms[position + 1..] {
286 let intersection: Vec<_> = left
287 .vertices
288 .iter()
289 .copied()
290 .filter(|vertex| right.vertices.binary_search(vertex).is_ok())
291 .collect();
292 if intersection.len() < 2 {
293 continue;
294 }
295 for (position, &u) in intersection.iter().enumerate() {
296 for &v in &intersection[position + 1..] {
297 edges.insert(EdgeKey::new(u, v));
298 }
299 }
300 }
301 }
302 edges.into_iter().collect()
303}
304
305pub(crate) fn compose_result(
306 input: &SparseDistanceMatrix,
307 params: &RipsParams,
308 states: &[ProgramAtomState],
309) -> Result<ExplainedDiagram> {
310 let mut seeds = Vec::new();
311 let terminal = terminal_level(input, params.threshold);
312 for state in states {
313 for space in &state.explained.spaces {
314 let interval = space.interval;
315 let scale = if interval.death.is_finite() {
316 previous_float(interval.death)
317 } else {
318 terminal
319 };
320 let cocycles = space
321 .basis
322 .iter()
323 .map(|class| Cocycle {
324 modulus: class.cocycle.modulus,
325 scale,
326 terms: class
327 .cocycle
328 .terms
329 .iter()
330 .map(|term| CocycleTerm {
331 u: state.vertices[term.u],
332 v: state.vertices[term.v],
333 coefficient: term.coefficient,
334 })
335 .collect(),
336 })
337 .collect();
338 let critical_pairs = space
339 .critical_pairs
340 .iter()
341 .map(|pair| map_critical_pair(pair, &state.vertices))
342 .collect();
343 seeds.push(SpaceSeed {
344 interval,
345 cocycles,
346 critical_pairs,
347 });
348 }
349 }
350 let spaces = merge_spaces(input, params.modulus, seeds)?;
351 let (mut diagram, _) = h0_diagram(input, params.threshold);
352 for space in &spaces {
353 diagram
354 .bars
355 .extend(std::iter::repeat_n(space.interval, space.basis.len()));
356 }
357 diagram.canonicalize();
358 Ok(ExplainedDiagram { diagram, spaces })
359}
360
361#[derive(Debug)]
362struct SpaceSeed {
363 interval: Bar,
364 cocycles: Vec<Cocycle>,
365 critical_pairs: Vec<CriticalPair>,
366}
367
368fn merge_spaces(
369 input: &SparseDistanceMatrix,
370 modulus: u32,
371 seeds: Vec<SpaceSeed>,
372) -> Result<Vec<PersistentClassSpace>> {
373 let mut groups: BTreeMap<(u64, u64), SpaceSeed> = BTreeMap::new();
374 for seed in seeds {
375 let key = (seed.interval.birth.to_bits(), seed.interval.death.to_bits());
376 let group = groups.entry(key).or_insert_with(|| SpaceSeed {
377 interval: seed.interval,
378 cocycles: Vec::new(),
379 critical_pairs: Vec::new(),
380 });
381 group.cocycles.extend(seed.cocycles);
382 group.critical_pairs.extend(seed.critical_pairs);
383 }
384 let mut spaces = Vec::with_capacity(groups.len());
385 for mut seed in groups.into_values() {
386 let cocycles = canonical_space_basis(input, modulus, &seed.cocycles)?;
387 if cocycles.len() != seed.critical_pairs.len() {
388 return Err(Error::InvalidInput(format!(
389 "composed class-space rank {} differs from critical-pair count {}",
390 cocycles.len(),
391 seed.critical_pairs.len()
392 )));
393 }
394 for cocycle in &cocycles {
395 validate_h1_cocycle(input, cocycle)?;
396 }
397 seed.critical_pairs.sort_by(critical_pair_order);
398 let id = group_id(seed.interval, modulus, &cocycles);
399 let basis = cocycles
400 .into_iter()
401 .enumerate()
402 .map(|(basis_index, cocycle)| PersistentClass {
403 id: basis_class_id(id, basis_index, &cocycle),
404 group_id: id,
405 basis_index,
406 interval: seed.interval,
407 cocycle,
408 provenance: None,
409 })
410 .collect();
411 spaces.push(PersistentClassSpace {
412 id,
413 interval: seed.interval,
414 basis,
415 critical_pairs: seed.critical_pairs,
416 });
417 }
418 spaces.sort_by(|a, b| {
419 a.interval
420 .birth
421 .total_cmp(&b.interval.birth)
422 .then(a.interval.death.total_cmp(&b.interval.death))
423 .then(a.id.cmp(&b.id))
424 });
425 Ok(spaces)
426}
427
428pub(super) fn reweight_explained(
429 input: &SparseDistanceMatrix,
430 previous: &ExplainedDiagram,
431 evaluation: &crate::CertifiedRegionEvaluation,
432 modulus: u32,
433 threshold: Option<f64>,
434) -> Result<Option<ExplainedDiagram>> {
435 if evaluation.h1_critical_pairs().len() != previous.class_count() {
436 return Ok(None);
437 }
438 let records: BTreeMap<_, _> = evaluation
439 .h1_critical_pairs()
440 .iter()
441 .map(|(bar, pair)| (critical_pair_key(pair), (*bar, pair.clone())))
442 .collect();
443 let terminal = terminal_level(input, threshold);
444 let mut seeds = Vec::new();
445 for space in &previous.spaces {
446 let mut interval = None;
447 let mut pairs = Vec::with_capacity(space.critical_pairs.len());
448 for pair in &space.critical_pairs {
449 let Some((bar, updated_pair)) = records.get(&critical_pair_key(pair)) else {
450 return Ok(None);
451 };
452 if interval.is_some_and(|interval: Bar| !bar_bits_equal(interval, *bar)) {
453 return Ok(None);
454 }
455 interval = Some(*bar);
456 pairs.push(updated_pair.clone());
457 }
458 let Some(interval) = interval else {
459 return Ok(None);
460 };
461 let scale = if interval.death.is_finite() {
462 previous_float(interval.death)
463 } else {
464 terminal
465 };
466 let cocycles: Vec<_> = space
467 .basis
468 .iter()
469 .map(|class| Cocycle {
470 modulus,
471 scale,
472 terms: class.cocycle.terms.clone(),
473 })
474 .collect();
475 if cocycles
476 .iter()
477 .any(|cocycle| validate_h1_cocycle(input, cocycle).is_err())
478 {
479 return Ok(None);
480 }
481 seeds.push(SpaceSeed {
482 interval,
483 cocycles,
484 critical_pairs: pairs,
485 });
486 }
487 let spaces = merge_spaces(input, modulus, seeds)?;
488 Ok(Some(ExplainedDiagram {
489 diagram: evaluation.diagram().clone(),
490 spaces,
491 }))
492}