Skip to main content

holos_tda/program/
diagram.rs

1use crate::{Error, ExplainedDiagram, RegionViolationKind, Result, SparseDistanceMatrix};
2
3use super::composition::local_matrix;
4use super::model::{
5    CorrespondenceMode, PersistenceProgram, ProgramAtomState, ProgramDiagramState,
6    ProgramDiagramUpdate, ProgramDiagramUpdateMode, ProgramEvent, ProgramEventKind, ProgramWork,
7};
8use super::topology::{diagram_bits_equal, program_topology_events, simplex_edge};
9use super::update::recompile_work;
10
11impl PersistenceProgram {
12    /// Move a compiled program into stateful diagram evaluation.
13    pub fn into_diagram_state(self) -> ProgramDiagramState {
14        let graph = self.graph.clone();
15        let diagram = self.result.diagram.clone();
16        ProgramDiagramState {
17            program: self,
18            graph,
19            diagram,
20            dirty: false,
21        }
22    }
23}
24
25impl ProgramDiagramState {
26    /// Copy a compiled program into stateful diagram evaluation.
27    pub fn from_program(program: &PersistenceProgram) -> Self {
28        program.clone().into_diagram_state()
29    }
30
31    /// Current exact persistence diagram.
32    pub fn diagram(&self) -> &crate::Diagram {
33        &self.diagram
34    }
35
36    /// Advance the state and return only the exact persistence diagram.
37    pub fn advance(&mut self, updated: &SparseDistanceMatrix) -> Result<ProgramDiagramUpdate> {
38        let topology_events = program_topology_events(&self.program, updated);
39        if !topology_events.is_empty() {
40            return self.recompile(updated, topology_events);
41        }
42        if graph_bits_equal(&self.graph, updated) {
43            return Ok(ProgramDiagramUpdate {
44                diagram: self.diagram.clone(),
45                mode: ProgramDiagramUpdateMode::Reused,
46                events: Vec::new(),
47                work: ProgramWork {
48                    edges_checked: self.program.topology.len(),
49                    ..ProgramWork::default()
50                },
51            });
52        }
53
54        match self.program.evaluate_diagram(updated) {
55            Ok(evaluation) => {
56                let diagram = evaluation.diagram;
57                self.graph = updated.clone();
58                self.diagram = diagram.clone();
59                self.dirty = true;
60                Ok(ProgramDiagramUpdate {
61                    diagram,
62                    mode: ProgramDiagramUpdateMode::Reused,
63                    events: Vec::new(),
64                    work: evaluation.work,
65                })
66            }
67            Err(_) => {
68                let events = region_violation_events(&self.program, updated)?;
69                self.recompile(updated, events)
70            }
71        }
72    }
73
74    /// Materialize the current canonical class spaces and critical pairs.
75    pub fn materialize(&mut self) -> Result<&ExplainedDiagram> {
76        if !self.dirty {
77            return Ok(self.program.result());
78        }
79
80        let mut candidate = self.program.clone();
81        let update = match candidate.advance_with(&self.graph, CorrespondenceMode::Omit) {
82            Ok(update) => update,
83            Err(_) => return self.recompile_for_materialization(),
84        };
85        if !diagram_bits_equal(&update.result.diagram, &self.diagram) {
86            return self.recompile_for_materialization();
87        }
88
89        self.program = candidate;
90        self.dirty = false;
91        Ok(self.program.result())
92    }
93
94    /// Materialize and consume this state as a normal persistence program.
95    ///
96    /// This method consumes the state even when materialization returns an
97    /// error. Call [`ProgramDiagramState::materialize`] when the state must
98    /// remain available after a failed operation.
99    pub fn into_program(mut self) -> Result<PersistenceProgram> {
100        self.materialize()?;
101        Ok(self.program)
102    }
103
104    fn recompile(
105        &mut self,
106        updated: &SparseDistanceMatrix,
107        events: Vec<ProgramEvent>,
108    ) -> Result<ProgramDiagramUpdate> {
109        let replacement =
110            PersistenceProgram::compile(updated, &self.program.params, self.program.limits)?;
111        let work = recompile_work(self.program.topology.len(), &replacement);
112        let diagram = replacement.result().diagram.clone();
113        self.program = replacement;
114        self.graph = updated.clone();
115        self.diagram = diagram.clone();
116        self.dirty = false;
117        Ok(ProgramDiagramUpdate {
118            diagram,
119            mode: ProgramDiagramUpdateMode::Recompiled,
120            events,
121            work,
122        })
123    }
124
125    fn recompile_for_materialization(&mut self) -> Result<&ExplainedDiagram> {
126        let graph = self.graph.clone();
127        let params = self.program.params.clone();
128        let replacement = PersistenceProgram::compile(&graph, &params, self.program.limits)?;
129        let diagram = &replacement.result().diagram;
130        if !diagram_bits_equal(diagram, &self.diagram) {
131            return Err(Error::InvalidInput(
132                "diagram materialization fallback differs from the accepted diagram".into(),
133            ));
134        }
135        self.program = replacement;
136        self.dirty = false;
137        Ok(self.program.result())
138    }
139}
140
141fn graph_bits_equal(left: &SparseDistanceMatrix, right: &SparseDistanceMatrix) -> bool {
142    left.len() == right.len()
143        && left.num_edges() == right.num_edges()
144        && left
145            .edges()
146            .zip(right.edges())
147            .all(|((lu, lv, lw), (ru, rv, rw))| {
148                lu == ru && lv == rv && lw.to_bits() == rw.to_bits()
149            })
150}
151
152fn region_violation_events(
153    program: &PersistenceProgram,
154    updated: &SparseDistanceMatrix,
155) -> Result<Vec<ProgramEvent>> {
156    let mut events = Vec::new();
157    for state in program.states() {
158        let local = local_matrix(&state.vertices, &state.edges, updated)?;
159        events.extend(
160            state
161                .region
162                .violations(&local)
163                .iter()
164                .map(|violation| violation_event(state, violation)),
165        );
166    }
167    Ok(events)
168}
169
170fn violation_event(state: &ProgramAtomState, violation: &crate::RegionViolation) -> ProgramEvent {
171    let kind = match violation.kind() {
172        RegionViolationKind::VertexSetChanged => ProgramEventKind::VertexSetChanged,
173        RegionViolationKind::EdgeSetChanged => ProgramEventKind::EdgeSetChanged,
174        RegionViolationKind::ThresholdCrossing => ProgramEventKind::ThresholdCrossing,
175        RegionViolationKind::GuardFailed => ProgramEventKind::GuardFailed,
176    };
177    ProgramEvent {
178        kind,
179        atom: Some(state.info_index),
180        edge: violation.first().and_then(simplex_edge),
181        guard: violation
182            .guard_index()
183            .map(|index| state.region.guards()[index].kind()),
184    }
185}