Skip to main content

holos_tda/program/
composition.rs

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    /// Compile a compositional H0 and H1 program.
21    ///
22    /// Each cyclic atom receives its own reduction certificate. Graphs
23    /// without a useful split remain one exact atom.
24    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}