holos_tda/persistent_class_artifact/
build.rs1use crate::certificate::{CertificateError, CertificateLimits, ReductionCertificate};
2use crate::field::is_prime;
3use crate::{RipsParams, SparseDistanceMatrix, rips_persistence_with_classes_sparse};
4
5use super::model::{PersistenceCycleTerm, PersistenceTriangleTerm, PersistentClassArtifact};
6
7impl PersistentClassArtifact {
8 pub fn build(
16 input: &SparseDistanceMatrix,
17 params: &RipsParams,
18 space_index: usize,
19 basis_index: usize,
20 limits: CertificateLimits,
21 ) -> Result<Self, CertificateError> {
22 preflight(input, params, limits)?;
23 let mut class_params = params.clone();
24 class_params.max_dim = 1;
25 let explained = rips_persistence_with_classes_sparse(input, &class_params)
26 .map_err(|error| certificate_error(error.to_string()))?;
27 let space = explained.spaces.get(space_index).ok_or_else(|| {
28 certificate_error(format!(
29 "persistent class space {space_index} is out of range"
30 ))
31 })?;
32 let class = space.basis.get(basis_index).ok_or_else(|| {
33 certificate_error(format!(
34 "persistent class basis index {basis_index} is out of range"
35 ))
36 })?;
37 if space.critical_pairs.len() != space.basis.len() {
38 return Err(certificate_error(
39 "persistent class space has an incomplete critical-pair list",
40 ));
41 }
42 class
43 .validate_provenance(input)
44 .map_err(|error| certificate_error(error.to_string()))?;
45 let mut reduction_params = RipsParams::new(1).with_modulus(params.modulus);
46 reduction_params.threshold = params.threshold;
47 let witness = ReductionCertificate::build_cycle_witness(
48 input,
49 &reduction_params,
50 &space.critical_pairs,
51 &class.cocycle,
52 limits,
53 )?;
54 let mut cycle: Vec<_> = witness
55 .cycle
56 .into_iter()
57 .map(|(u, v, coefficient)| PersistenceCycleTerm { u, v, coefficient })
58 .collect();
59 cycle.sort_unstable();
60 let mut bounding_chain: Vec<_> = witness
61 .bounding_chain
62 .into_iter()
63 .map(|(vertices, coefficient)| PersistenceTriangleTerm {
64 vertices,
65 coefficient,
66 })
67 .collect();
68 bounding_chain.sort_unstable();
69 let artifact = Self {
70 source: input.clone(),
71 threshold: params.threshold,
72 class: class.clone(),
73 critical_pair: witness.pair,
74 cycle,
75 bounding_chain,
76 };
77 super::wire::validate_for_build(&artifact, limits)?;
78 Ok(artifact)
79 }
80}
81
82fn preflight(
83 input: &SparseDistanceMatrix,
84 params: &RipsParams,
85 limits: CertificateLimits,
86) -> Result<(), CertificateError> {
87 validate_parameters(params, limits)?;
88 validate_source_limits(input, limits)?;
89 validate_modulus(params.modulus)?;
90 validate_threshold(params.threshold)?;
91 let minimum = minimum_bytes(input.num_edges(), params.threshold.is_some())?;
92 if minimum > limits.max_bytes {
93 return Err(certificate_error(
94 "persistent class artifact exceeds its byte limit",
95 ));
96 }
97 let threshold = params.threshold.unwrap_or(f64::INFINITY);
98 count_triangles(input, threshold, limits.max_triangles)?;
99 Ok(())
100}
101
102fn validate_parameters(
103 params: &RipsParams,
104 limits: CertificateLimits,
105) -> Result<(), CertificateError> {
106 if params.max_dim < 1 {
107 return Err(certificate_error(
108 "persistent H1 classes require max_dim of at least 1",
109 ));
110 }
111 if params.collapse_edges {
112 return Err(certificate_error(
113 "persistent class artifacts do not accept edge collapse",
114 ));
115 }
116 if limits.max_dimension < 1 {
117 return Err(certificate_error(
118 "persistent H1 artifacts require a dimension limit of at least 1",
119 ));
120 }
121 if limits.max_bars == 0 {
122 return Err(certificate_error("persistent class bar exceeds its limit"));
123 }
124 if limits.max_terms < 2 {
125 return Err(certificate_error(
126 "persistent class cycle and cocycle exceed the term limit",
127 ));
128 }
129 Ok(())
130}
131
132fn validate_source_limits(
133 input: &SparseDistanceMatrix,
134 limits: CertificateLimits,
135) -> Result<(), CertificateError> {
136 if input.len() > limits.max_vertices {
137 return Err(certificate_error(format!(
138 "{} source vertices exceed the limit {}",
139 input.len(),
140 limits.max_vertices
141 )));
142 }
143 if input.num_edges() > limits.max_edges {
144 return Err(certificate_error(format!(
145 "{} source edges exceed the limit {}",
146 input.num_edges(),
147 limits.max_edges
148 )));
149 }
150 Ok(())
151}
152
153pub(super) fn count_triangles(
154 input: &SparseDistanceMatrix,
155 threshold: f64,
156 maximum: usize,
157) -> Result<usize, CertificateError> {
158 let upper = filtered_upper_neighbors(input, threshold);
159 let mut total = 0usize;
160 for u in 0..upper.len() {
161 for &v in &upper[u] {
162 let left = upper[u].partition_point(|&w| w <= v);
163 let right = upper[v].partition_point(|&w| w <= v);
164 count_common_neighbors(&upper[u][left..], &upper[v][right..], maximum, &mut total)?;
165 }
166 }
167 Ok(total)
168}
169
170fn filtered_upper_neighbors(input: &SparseDistanceMatrix, threshold: f64) -> Vec<Vec<usize>> {
171 let mut upper = vec![Vec::new(); input.len()];
172 for (u, v, value) in input.edges() {
173 if value <= threshold {
174 upper[u].push(v);
175 }
176 }
177 upper
178}
179
180fn count_common_neighbors(
181 left: &[usize],
182 right: &[usize],
183 maximum: usize,
184 total: &mut usize,
185) -> Result<(), CertificateError> {
186 let mut left_index = 0;
187 let mut right_index = 0;
188 while left_index < left.len() && right_index < right.len() {
189 match left[left_index].cmp(&right[right_index]) {
190 std::cmp::Ordering::Less => left_index += 1,
191 std::cmp::Ordering::Greater => right_index += 1,
192 std::cmp::Ordering::Equal => {
193 add_triangle(total, maximum)?;
194 left_index += 1;
195 right_index += 1;
196 }
197 }
198 }
199 Ok(())
200}
201
202fn add_triangle(total: &mut usize, maximum: usize) -> Result<(), CertificateError> {
203 *total = (*total)
204 .checked_add(1)
205 .ok_or_else(|| certificate_error("filtered triangle count overflows"))?;
206 if *total > maximum {
207 return Err(certificate_error(format!(
208 "filtered triangle count exceeds the limit {maximum}"
209 )));
210 }
211 Ok(())
212}
213
214fn validate_modulus(modulus: u32) -> Result<(), CertificateError> {
215 if modulus >= 32_768 || !is_prime(u64::from(modulus)) {
216 return Err(certificate_error(format!(
217 "modulus must be a prime below 32768, got {modulus}"
218 )));
219 }
220 Ok(())
221}
222
223fn validate_threshold(threshold: Option<f64>) -> Result<(), CertificateError> {
224 if threshold.is_some_and(|value| {
225 value.is_nan() || value < 0.0 || (value == 0.0 && value.to_bits() != 0)
226 }) {
227 return Err(certificate_error(
228 "threshold must be non-negative and not negative zero",
229 ));
230 }
231 Ok(())
232}
233
234fn minimum_bytes(edges: usize, has_threshold: bool) -> Result<usize, CertificateError> {
235 let fixed = 8usize
236 .checked_add(2)
237 .and_then(|value| value.checked_add(1))
238 .and_then(|value| value.checked_add(4))
239 .and_then(|value| value.checked_add(8 + 8))
240 .and_then(|value| value.checked_add(if has_threshold { 9 } else { 1 }))
241 .and_then(|value| value.checked_add(32 + 32 + 8 + 8 + 8 + 8))
242 .and_then(|value| value.checked_add(8 + 20))
243 .and_then(|value| value.checked_add(16 + 1))
244 .and_then(|value| value.checked_add(8 + 20))
245 .and_then(|value| value.checked_add(8))
246 .and_then(|value| value.checked_add(32))
247 .ok_or_else(|| certificate_error("persistent class artifact byte count overflows"))?;
248 let source = edges
249 .checked_mul(24)
250 .ok_or_else(|| certificate_error("persistent class artifact byte count overflows"))?;
251 fixed
252 .checked_add(source)
253 .ok_or_else(|| certificate_error("persistent class artifact byte count overflows"))
254}
255
256fn certificate_error(message: impl Into<String>) -> CertificateError {
257 CertificateError::new(message)
258}