Skip to main content

holos_tda/program/
update.rs

1use std::collections::BTreeSet;
2
3use crate::{
4    CertificateLimits, ClassCorrespondence, EdgeKey, Error, ExplainedDiagram, PersistentClassSpace,
5    Result, RipsParams, SparseDistanceMatrix,
6};
7
8use super::composition::{compose_result, local_matrix, reweight_explained};
9use super::continuation::class_continuation;
10use super::model::{
11    CorrespondenceMode, PersistenceProgram, ProgramAtomInfo, ProgramAtomState, ProgramBranch,
12    ProgramCheckpoint, ProgramEvaluation, ProgramEvent, ProgramSummary, ProgramUpdate,
13    ProgramUpdateMode, ProgramWork,
14};
15use super::topology::{check_program_topology, h0_diagram, h0_provenance, program_topology_events};
16
17mod atom;
18
19use atom::{preview_atom_state, update_atom_state};
20
21impl PersistenceProgram {
22    /// Exact current result.
23    pub fn result(&self) -> &ExplainedDiagram {
24        &self.result
25    }
26
27    /// Program decomposition and guard counts.
28    pub fn summary(&self) -> ProgramSummary {
29        self.summary
30    }
31
32    /// All structural atoms, including bridge atoms.
33    pub fn atoms(&self) -> &[ProgramAtomInfo] {
34        &self.atoms
35    }
36
37    /// Capture the current state for later restore or branching.
38    pub fn checkpoint(&self) -> ProgramCheckpoint {
39        ProgramCheckpoint {
40            program: self.clone(),
41        }
42    }
43
44    /// Replace this program with a captured state.
45    pub fn restore(&mut self, checkpoint: &ProgramCheckpoint) {
46        *self = checkpoint.program.clone();
47    }
48
49    /// Advance an ordered batch atomically.
50    ///
51    /// If one update fails, the program is left unchanged.
52    pub fn advance_batch(
53        &mut self,
54        updates: &[SparseDistanceMatrix],
55    ) -> Result<Vec<ProgramUpdate>> {
56        self.advance_batch_with(updates, CorrespondenceMode::Exact)
57    }
58
59    /// Advance an ordered batch atomically with correspondence control.
60    pub fn advance_batch_with(
61        &mut self,
62        updates: &[SparseDistanceMatrix],
63        correspondence_mode: CorrespondenceMode,
64    ) -> Result<Vec<ProgramUpdate>> {
65        let mut candidate = self.clone();
66        let mut results = Vec::with_capacity(updates.len());
67        for update in updates {
68            results.push(candidate.advance_with(update, correspondence_mode)?);
69        }
70        *self = candidate;
71        Ok(results)
72    }
73
74    /// Advance independent alternatives from the current state.
75    ///
76    /// Output order matches input order. The program is left unchanged. With
77    /// more than one configured thread, alternatives run concurrently.
78    pub fn branch(&self, alternatives: &[SparseDistanceMatrix]) -> Result<Vec<ProgramBranch>> {
79        self.branch_with(alternatives, CorrespondenceMode::Exact)
80    }
81
82    /// Advance independent alternatives with correspondence control.
83    pub fn branch_with(
84        &self,
85        alternatives: &[SparseDistanceMatrix],
86        correspondence_mode: CorrespondenceMode,
87    ) -> Result<Vec<ProgramBranch>> {
88        if alternatives.is_empty() {
89            return Ok(Vec::new());
90        }
91        if self.params.threads <= 1 || alternatives.len() == 1 {
92            return alternatives
93                .iter()
94                .enumerate()
95                .map(|(index, alternative)| {
96                    let mut program = self.clone();
97                    let update = program.advance_with(alternative, correspondence_mode)?;
98                    Ok(ProgramBranch {
99                        index,
100                        update,
101                        program,
102                    })
103                })
104                .collect();
105        }
106
107        use rayon::prelude::*;
108        let workers = self.params.threads.min(alternatives.len());
109        let pool = rayon::ThreadPoolBuilder::new()
110            .num_threads(workers)
111            .build()
112            .map_err(|error| {
113                Error::InvalidInput(format!("cannot create branch workers: {error}"))
114            })?;
115        let results = pool.install(|| {
116            alternatives
117                .par_iter()
118                .enumerate()
119                .map(|(index, alternative)| {
120                    let mut program = self.clone();
121                    let update = program.advance_with(alternative, correspondence_mode)?;
122                    Ok(ProgramBranch {
123                        index,
124                        update,
125                        program,
126                    })
127                })
128                .collect::<Vec<Result<ProgramBranch>>>()
129        });
130        results.into_iter().collect()
131    }
132
133    /// Evaluate only the diagram while every touched certificate remains
134    /// valid.
135    pub fn evaluate_diagram(&self, updated: &SparseDistanceMatrix) -> Result<ProgramEvaluation> {
136        let mut work = ProgramWork {
137            edges_checked: self.topology.len(),
138            ..ProgramWork::default()
139        };
140        let edge_values = check_program_topology(self, updated)?;
141        let mut h1 = Vec::new();
142        for state in &self.states {
143            h1.extend(
144                state
145                    .region
146                    .evaluate_h1_indexed(&edge_values, &state.edge_positions)?,
147            );
148            work.guards_checked += state.region.guards().len();
149        }
150        let (mut diagram, scanned) = h0_diagram(updated, self.params.threshold);
151        work.h0_edges_scanned = scanned;
152        diagram.bars.extend(h1);
153        diagram.canonicalize();
154        Ok(ProgramEvaluation { diagram, work })
155    }
156
157    /// Advance to new weights with exact class correspondence.
158    pub fn advance(&mut self, updated: &SparseDistanceMatrix) -> Result<ProgramUpdate> {
159        self.advance_with(updated, CorrespondenceMode::Exact)
160    }
161
162    /// Advance with explicit control over cross-state class correspondence.
163    pub fn advance_with(
164        &mut self,
165        updated: &SparseDistanceMatrix,
166        correspondence_mode: CorrespondenceMode,
167    ) -> Result<ProgramUpdate> {
168        let old = self.result.clone();
169        let old_graph = self.graph.clone();
170        let topology_events = program_topology_events(self, updated);
171        if !topology_events.is_empty() {
172            return self.recompile_update(
173                updated,
174                old,
175                old_graph,
176                correspondence_mode,
177                topology_events,
178            );
179        }
180
181        let changed = changed_edges(self, updated);
182        let mut work = ProgramWork {
183            edges_checked: self.topology.len(),
184            ..ProgramWork::default()
185        };
186        let mut events = Vec::new();
187        for state in &mut self.states {
188            if !state.edges.iter().any(|edge| changed.contains(edge)) {
189                continue;
190            }
191            update_atom_state(
192                state,
193                updated,
194                &changed,
195                self.params.modulus,
196                self.limits,
197                &mut work,
198                &mut events,
199            )?;
200        }
201        self.finish_weight_update(updated, old, old_graph, correspondence_mode, events, work)
202    }
203
204    fn recompile_update(
205        &mut self,
206        updated: &SparseDistanceMatrix,
207        old: ExplainedDiagram,
208        old_graph: SparseDistanceMatrix,
209        correspondence_mode: CorrespondenceMode,
210        events: Vec<ProgramEvent>,
211    ) -> Result<ProgramUpdate> {
212        let replacement = Self::compile(updated, &self.params, self.limits)?;
213        let result = replacement.result.clone();
214        let correspondence = update_correspondence(
215            correspondence_mode,
216            &old_graph,
217            &old.spaces,
218            updated,
219            &result.spaces,
220            self.params.modulus,
221        )?;
222        let work = recompile_work(self.topology.len(), &replacement);
223        let continuation = class_continuation(&old.spaces, &result.spaces);
224        *self = replacement;
225        Ok(ProgramUpdate {
226            result,
227            mode: ProgramUpdateMode::Recompiled,
228            events,
229            continuation,
230            correspondence,
231            work,
232        })
233    }
234
235    fn finish_weight_update(
236        &mut self,
237        updated: &SparseDistanceMatrix,
238        old: ExplainedDiagram,
239        old_graph: SparseDistanceMatrix,
240        correspondence_mode: CorrespondenceMode,
241        events: Vec<ProgramEvent>,
242        mut work: ProgramWork,
243    ) -> Result<ProgramUpdate> {
244        self.graph = updated.clone();
245        self.result = compose_result(updated, &self.params, &self.states)?;
246        self.summary.guards = self
247            .states
248            .iter()
249            .map(|state| state.region.guards().len())
250            .sum();
251        let (h0_deaths, h0_essential, scanned) = h0_provenance(updated, self.params.threshold);
252        self.h0_deaths = h0_deaths;
253        self.h0_essential = h0_essential;
254        work.h0_edges_scanned = scanned;
255        let continuation = class_continuation(&old.spaces, &self.result.spaces);
256        let correspondence = update_correspondence(
257            correspondence_mode,
258            &old_graph,
259            &old.spaces,
260            updated,
261            &self.result.spaces,
262            self.params.modulus,
263        )?;
264        Ok(ProgramUpdate {
265            result: self.result.clone(),
266            mode: update_mode(&work),
267            events,
268            continuation,
269            correspondence,
270            work,
271        })
272    }
273
274    pub(crate) fn advance_reused(
275        &mut self,
276        updated: &SparseDistanceMatrix,
277    ) -> Result<ProgramUpdate> {
278        check_program_topology(self, updated)?;
279        let old = self.result.clone();
280        let old_graph = self.graph.clone();
281        let changed = changed_edges(self, updated);
282        let mut work = ProgramWork {
283            edges_checked: self.topology.len(),
284            ..ProgramWork::default()
285        };
286        for state in &mut self.states {
287            if !state.edges.iter().any(|edge| changed.contains(edge)) {
288                continue;
289            }
290            work.atoms_touched += 1;
291            work.guards_checked += state.region.guards().len();
292            let local = local_matrix(&state.vertices, &state.edges, updated)?;
293            let evaluation = state.region.evaluate(&local)?;
294            let Some(explained) = reweight_explained(
295                &local,
296                &state.explained,
297                &evaluation,
298                self.params.modulus,
299                self.params.threshold,
300            )?
301            else {
302                return Err(Error::InvalidInput(
303                    "program class-space state requires a checked checkpoint".into(),
304                ));
305            };
306            state.artifact = state
307                .artifact
308                .rebind(
309                    &state.certified_graph,
310                    &local,
311                    explained.clone(),
312                    self.limits,
313                )
314                .map_err(|error| Error::InvalidInput(error.to_string()))?;
315            state.certified_graph = local;
316            state.explained = explained;
317            work.atoms_reused += 1;
318        }
319        self.graph = updated.clone();
320        self.result = compose_result(updated, &self.params, &self.states)?;
321        let (h0_deaths, h0_essential, scanned) = h0_provenance(updated, self.params.threshold);
322        self.h0_deaths = h0_deaths;
323        self.h0_essential = h0_essential;
324        work.h0_edges_scanned = scanned;
325        Ok(ProgramUpdate {
326            result: self.result.clone(),
327            mode: ProgramUpdateMode::Reused,
328            events: Vec::new(),
329            continuation: class_continuation(&old.spaces, &self.result.spaces),
330            correspondence: crate::class_correspondences(
331                &old_graph,
332                &old.spaces,
333                updated,
334                &self.result.spaces,
335                self.params.modulus,
336            )?,
337            work,
338        })
339    }
340
341    pub(crate) fn preview_update(
342        &self,
343        updated: &SparseDistanceMatrix,
344        replacement_cyclic_atoms: usize,
345    ) -> Result<(ProgramUpdateMode, Vec<ProgramEvent>, ProgramWork)> {
346        let topology_events = program_topology_events(self, updated);
347        if !topology_events.is_empty() {
348            return Ok((
349                ProgramUpdateMode::Recompiled,
350                topology_events,
351                ProgramWork {
352                    edges_checked: self.topology.len().max(updated.num_edges()),
353                    h0_edges_scanned: updated
354                        .edges()
355                        .filter(|&(_, _, value)| {
356                            value <= self.params.threshold.unwrap_or(f64::INFINITY)
357                        })
358                        .count(),
359                    guards_checked: 0,
360                    atoms_touched: replacement_cyclic_atoms,
361                    atoms_reused: 0,
362                    atoms_repaired: 0,
363                    atoms_rebuilt: replacement_cyclic_atoms,
364                    reduction_columns_reused: 0,
365                    reduction_columns_reduced: 0,
366                    reduction_column_additions: 0,
367                },
368            ));
369        }
370        let changed = changed_edges(self, updated);
371        let mut events = Vec::new();
372        let mut work = ProgramWork {
373            edges_checked: self.topology.len(),
374            ..ProgramWork::default()
375        };
376        for state in &self.states {
377            if !state.edges.iter().any(|edge| changed.contains(edge)) {
378                continue;
379            }
380            preview_atom_state(
381                state,
382                updated,
383                &changed,
384                self.params.modulus,
385                self.limits,
386                &mut work,
387                &mut events,
388            )?;
389        }
390        work.h0_edges_scanned = updated
391            .edges()
392            .filter(|&(_, _, value)| value <= self.params.threshold.unwrap_or(f64::INFINITY))
393            .count();
394        Ok((update_mode(&work), events, work))
395    }
396
397    pub(crate) fn states(&self) -> &[ProgramAtomState] {
398        &self.states
399    }
400
401    pub(crate) fn params(&self) -> &RipsParams {
402        &self.params
403    }
404
405    pub(crate) fn limits(&self) -> CertificateLimits {
406        self.limits
407    }
408
409    pub(crate) fn current_graph(&self) -> &SparseDistanceMatrix {
410        &self.graph
411    }
412}
413
414fn changed_edges(
415    program: &PersistenceProgram,
416    updated: &SparseDistanceMatrix,
417) -> BTreeSet<EdgeKey> {
418    program
419        .topology
420        .iter()
421        .copied()
422        .filter(|edge| {
423            program.graph.get(edge.u, edge.v).to_bits() != updated.get(edge.u, edge.v).to_bits()
424        })
425        .collect()
426}
427
428fn update_correspondence(
429    mode: CorrespondenceMode,
430    old_graph: &SparseDistanceMatrix,
431    old_spaces: &[PersistentClassSpace],
432    updated: &SparseDistanceMatrix,
433    new_spaces: &[PersistentClassSpace],
434    modulus: u32,
435) -> Result<Vec<ClassCorrespondence>> {
436    if mode == CorrespondenceMode::Omit {
437        return Ok(Vec::new());
438    }
439    crate::class_correspondences(old_graph, old_spaces, updated, new_spaces, modulus)
440}
441
442pub(super) fn recompile_work(old_edges: usize, replacement: &PersistenceProgram) -> ProgramWork {
443    let threshold = replacement.params.threshold.unwrap_or(f64::INFINITY);
444    ProgramWork {
445        edges_checked: old_edges.max(replacement.topology.len()),
446        h0_edges_scanned: replacement
447            .graph
448            .edges()
449            .filter(|&(_, _, value)| value <= threshold)
450            .count(),
451        atoms_touched: replacement.states.len(),
452        atoms_rebuilt: replacement.states.len(),
453        reduction_columns_reduced: replacement
454            .states
455            .iter()
456            .map(|state| {
457                let certificate = state.artifact.reduction_certificate();
458                certificate.edge_columns().len() + certificate.triangle_columns().len()
459            })
460            .sum(),
461        ..ProgramWork::default()
462    }
463}
464
465fn update_mode(work: &ProgramWork) -> ProgramUpdateMode {
466    if work.atoms_rebuilt == 0 && work.atoms_repaired == 0 {
467        ProgramUpdateMode::Reused
468    } else {
469        ProgramUpdateMode::Repaired
470    }
471}