holos_tda/program/
diagram.rs1use 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 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 pub fn from_program(program: &PersistenceProgram) -> Self {
28 program.clone().into_diagram_state()
29 }
30
31 pub fn diagram(&self) -> &crate::Diagram {
33 &self.diagram
34 }
35
36 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 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 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, ¶ms, 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}