Skip to main content

holos_tda/
zigzag.rs

1//! Exact interval decomposition of finite zigzag vector-space modules.
2//!
3//! A module is a type-A quiver over one declared prime field. Arrows may
4//! point in either direction. The decomposition uses the generalized rank of
5//! every contiguous submodule and Möbius inversion. Repeated interval
6//! summands remain one space with a multiplicity.
7
8mod algebra;
9mod ranks;
10use ranks::generalized_rank;
11use std::fmt;
12
13use sha2::{Digest, Sha256};
14
15use crate::field::{MODULUS_LIMIT, is_prime};
16use crate::{Error, Result};
17
18/// Resource limits for exact zigzag decomposition.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20#[non_exhaustive]
21pub struct ZigzagLimits {
22    /// Largest accepted node count.
23    pub max_nodes: usize,
24    /// Largest sum of node dimensions.
25    pub max_total_dimension: usize,
26    /// Largest total nonzero map coefficient count.
27    pub max_map_terms: usize,
28    /// Largest generalized-rank work count.
29    pub max_rank_work: usize,
30}
31
32impl Default for ZigzagLimits {
33    fn default() -> Self {
34        Self {
35            max_nodes: 2_049,
36            max_total_dimension: 100_000,
37            max_map_terms: 20_000_000,
38            max_rank_work: 100_000_000,
39        }
40    }
41}
42
43/// Direction of one arrow between adjacent module nodes.
44#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
45pub enum ZigzagDirection {
46    /// The map goes from the left node to the right node.
47    Forward,
48    /// The map goes from the right node to the left node.
49    Backward,
50}
51
52/// One nonzero coefficient in a zigzag map column.
53#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
54pub struct ZigzagTerm {
55    /// Target basis position.
56    pub target: usize,
57    /// Coefficient in `1..modulus`.
58    pub coefficient: u32,
59}
60
61/// One linear map between adjacent zigzag nodes.
62#[derive(Debug, Clone, PartialEq, Eq)]
63pub struct ZigzagMap {
64    direction: ZigzagDirection,
65    columns: Vec<Vec<ZigzagTerm>>,
66}
67
68impl ZigzagMap {
69    /// Construct a map from columns in source basis order.
70    pub fn new(direction: ZigzagDirection, columns: Vec<Vec<ZigzagTerm>>) -> Self {
71        Self { direction, columns }
72    }
73
74    /// Arrow direction.
75    pub fn direction(&self) -> ZigzagDirection {
76        self.direction
77    }
78
79    /// Map columns in source basis order.
80    pub fn columns(&self) -> &[Vec<ZigzagTerm>] {
81        &self.columns
82    }
83}
84
85/// Content identifier of one finite zigzag module.
86#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
87pub struct ZigzagModuleId([u8; 32]);
88
89impl ZigzagModuleId {
90    /// Raw identifier bytes.
91    pub fn as_bytes(&self) -> &[u8; 32] {
92        &self.0
93    }
94}
95
96impl fmt::Display for ZigzagModuleId {
97    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
98        write_hex(formatter, &self.0)
99    }
100}
101
102/// Content identifier of an interval-isotypic class space.
103#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
104pub struct ZigzagIntervalId([u8; 32]);
105
106impl ZigzagIntervalId {
107    /// Raw identifier bytes.
108    pub fn as_bytes(&self) -> &[u8; 32] {
109        &self.0
110    }
111}
112
113impl fmt::Display for ZigzagIntervalId {
114    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
115        write_hex(formatter, &self.0)
116    }
117}
118
119/// One interval-isotypic space in a zigzag decomposition.
120#[derive(Debug, Clone, PartialEq, Eq)]
121pub struct ZigzagInterval {
122    /// Content identifier of this interval and its source module.
123    pub id: ZigzagIntervalId,
124    /// First node covered by the interval.
125    pub start: usize,
126    /// Last node covered by the interval, inclusive.
127    pub end: usize,
128    /// Number of indistinguishable copies of this interval summand.
129    pub multiplicity: usize,
130}
131
132/// Exact interval decomposition and generalized ranks of a finite zigzag.
133#[derive(Debug, Clone, PartialEq, Eq)]
134pub struct ZigzagBarcode {
135    /// Content identifier of the checked module.
136    pub module: ZigzagModuleId,
137    /// Node dimensions in zigzag order.
138    pub dimensions: Vec<usize>,
139    /// Generalized rank for every `[start, end]`, stored row-major.
140    pub generalized_ranks: Vec<usize>,
141    /// Nonzero interval multiplicities in lexicographic endpoint order.
142    pub intervals: Vec<ZigzagInterval>,
143}
144
145impl ZigzagBarcode {
146    /// Generalized rank on the inclusive subinterval `[start, end]`.
147    pub fn rank(&self, start: usize, end: usize) -> Option<usize> {
148        (start <= end && end < self.dimensions.len())
149            .then(|| self.generalized_ranks[start * self.dimensions.len() + end])
150    }
151}
152
153/// A checked finite zigzag vector-space module over a prime field.
154#[derive(Debug, Clone, PartialEq, Eq)]
155pub struct ZigzagModule {
156    id: ZigzagModuleId,
157    modulus: u32,
158    dimensions: Vec<usize>,
159    maps: Vec<ZigzagMap>,
160    limits: ZigzagLimits,
161}
162
163impl ZigzagModule {
164    /// Construct and validate one finite zigzag module.
165    pub fn new(
166        modulus: u32,
167        dimensions: Vec<usize>,
168        maps: Vec<ZigzagMap>,
169        limits: ZigzagLimits,
170    ) -> Result<Self> {
171        validate_module_shape(modulus, &dimensions, &maps, limits)?;
172        let total_dimension = total_dimension(&dimensions)?;
173        if total_dimension > limits.max_total_dimension {
174            return Err(Error::InvalidInput(format!(
175                "zigzag total dimension exceeds the limit {}",
176                limits.max_total_dimension
177            )));
178        }
179        let mut terms = 0usize;
180        for (position, map) in maps.iter().enumerate() {
181            terms = terms
182                .checked_add(validate_map(position, map, &dimensions, modulus)?)
183                .ok_or_else(|| Error::InvalidInput("zigzag map term count overflows".into()))?;
184        }
185        if terms > limits.max_map_terms {
186            return Err(Error::InvalidInput(format!(
187                "zigzag map term count exceeds the limit {}",
188                limits.max_map_terms
189            )));
190        }
191        let id = module_id(modulus, &dimensions, &maps);
192        Ok(Self {
193            id,
194            modulus,
195            dimensions,
196            maps,
197            limits,
198        })
199    }
200
201    /// Content identifier of this module.
202    pub fn id(&self) -> ZigzagModuleId {
203        self.id
204    }
205
206    /// Prime coefficient modulus.
207    pub fn modulus(&self) -> u32 {
208        self.modulus
209    }
210
211    /// Node dimensions in zigzag order.
212    pub fn dimensions(&self) -> &[usize] {
213        &self.dimensions
214    }
215
216    /// Adjacent maps in zigzag order.
217    pub fn maps(&self) -> &[ZigzagMap] {
218        &self.maps
219    }
220
221    /// Decompose this type-A representation into interval summands.
222    pub fn decompose(&self) -> Result<ZigzagBarcode> {
223        let generalized_ranks = self.compute_generalized_ranks()?;
224        let intervals =
225            interval_multiplicities(self.id, &generalized_ranks, self.dimensions.len())?;
226        Ok(ZigzagBarcode {
227            module: self.id,
228            dimensions: self.dimensions.clone(),
229            generalized_ranks,
230            intervals,
231        })
232    }
233
234    fn compute_generalized_ranks(&self) -> Result<Vec<usize>> {
235        let count = self.dimensions.len();
236        let mut ranks = vec![0usize; count * count];
237        let mut work = 0usize;
238        for start in (0..count).rev() {
239            for end in start..count {
240                work = work
241                    .checked_add(self.rank_work(start, end)?)
242                    .ok_or_else(rank_work_overflow)?;
243                if work > self.limits.max_rank_work {
244                    return Err(Error::InvalidInput(format!(
245                        "zigzag generalized-rank work exceeds the limit {}",
246                        self.limits.max_rank_work
247                    )));
248                }
249                ranks[start * count + end] = generalized_rank(self, start, end)?;
250            }
251        }
252        Ok(ranks)
253    }
254
255    fn rank_work(&self, start: usize, end: usize) -> Result<usize> {
256        let ambient = self.dimensions[start..=end].iter().sum::<usize>();
257        let arrows = (start..end).try_fold(0usize, |sum, position| {
258            let (source, target) =
259                map_shape(&self.dimensions, position, self.maps[position].direction);
260            sum.checked_add(source)
261                .and_then(|value| value.checked_add(target))
262                .ok_or_else(rank_work_overflow)
263        })?;
264        ambient.checked_add(arrows).ok_or_else(rank_work_overflow)
265    }
266}
267
268fn validate_module_shape(
269    modulus: u32,
270    dimensions: &[usize],
271    maps: &[ZigzagMap],
272    limits: ZigzagLimits,
273) -> Result<()> {
274    if !is_prime(modulus as u64) || u64::from(modulus) >= MODULUS_LIMIT {
275        return Err(Error::InvalidInput(
276            "zigzag modulus must be a supported prime".into(),
277        ));
278    }
279    if dimensions.is_empty() || dimensions.len() > limits.max_nodes {
280        return Err(Error::InvalidInput(format!(
281            "zigzag node count must be in 1..={}",
282            limits.max_nodes
283        )));
284    }
285    if maps.len() + 1 != dimensions.len() {
286        return Err(Error::InvalidInput(
287            "zigzag requires one map between each adjacent node".into(),
288        ));
289    }
290    Ok(())
291}
292
293fn total_dimension(dimensions: &[usize]) -> Result<usize> {
294    dimensions.iter().try_fold(0usize, |sum, value| {
295        sum.checked_add(*value)
296            .ok_or_else(|| Error::InvalidInput("zigzag total dimension overflows".into()))
297    })
298}
299
300fn validate_map(
301    position: usize,
302    map: &ZigzagMap,
303    dimensions: &[usize],
304    modulus: u32,
305) -> Result<usize> {
306    let (source, target) = map_shape(dimensions, position, map.direction);
307    if map.columns.len() != source {
308        return Err(Error::InvalidInput(format!(
309            "zigzag map {position} has {} columns but its source dimension is {source}",
310            map.columns.len()
311        )));
312    }
313    let mut terms = 0usize;
314    for column in &map.columns {
315        validate_column(position, column, target, modulus)?;
316        terms = terms
317            .checked_add(column.len())
318            .ok_or_else(|| Error::InvalidInput("zigzag map term count overflows".into()))?;
319    }
320    Ok(terms)
321}
322
323fn validate_column(
324    position: usize,
325    column: &[ZigzagTerm],
326    target: usize,
327    modulus: u32,
328) -> Result<()> {
329    let mut previous = None;
330    for term in column {
331        if term.target >= target
332            || term.coefficient == 0
333            || term.coefficient >= modulus
334            || previous.is_some_and(|value| value >= term.target)
335        {
336            return Err(Error::InvalidInput(format!(
337                "zigzag map {position} has a noncanonical term"
338            )));
339        }
340        previous = Some(term.target);
341    }
342    Ok(())
343}
344
345fn rank_work_overflow() -> Error {
346    Error::InvalidInput("zigzag rank work overflows".into())
347}
348
349fn interval_multiplicities(
350    module: ZigzagModuleId,
351    ranks: &[usize],
352    count: usize,
353) -> Result<Vec<ZigzagInterval>> {
354    let mut intervals = Vec::new();
355    for start in 0..count {
356        for end in start..count {
357            if let Some(interval) = interval_multiplicity(module, ranks, count, start, end)? {
358                intervals.push(interval);
359            }
360        }
361    }
362    Ok(intervals)
363}
364
365fn interval_multiplicity(
366    module: ZigzagModuleId,
367    ranks: &[usize],
368    count: usize,
369    start: usize,
370    end: usize,
371) -> Result<Option<ZigzagInterval>> {
372    let rank = |left: usize, right: usize| -> i128 { ranks[left * count + right] as i128 };
373    let mut multiplicity = rank(start, end);
374    if start > 0 {
375        multiplicity -= rank(start - 1, end);
376    }
377    if end + 1 < count {
378        multiplicity -= rank(start, end + 1);
379    }
380    if start > 0 && end + 1 < count {
381        multiplicity += rank(start - 1, end + 1);
382    }
383    if multiplicity < 0 {
384        return Err(Error::InvalidInput(
385            "zigzag generalized ranks violate interval decomposability".into(),
386        ));
387    }
388    if multiplicity == 0 {
389        return Ok(None);
390    }
391    let multiplicity = usize::try_from(multiplicity)
392        .map_err(|_| Error::InvalidInput("zigzag interval multiplicity overflows".into()))?;
393    Ok(Some(ZigzagInterval {
394        id: interval_id(module, start, end),
395        start,
396        end,
397        multiplicity,
398    }))
399}
400
401fn map_shape(dimensions: &[usize], position: usize, direction: ZigzagDirection) -> (usize, usize) {
402    match direction {
403        ZigzagDirection::Forward => (dimensions[position], dimensions[position + 1]),
404        ZigzagDirection::Backward => (dimensions[position + 1], dimensions[position]),
405    }
406}
407
408fn module_id(modulus: u32, dimensions: &[usize], maps: &[ZigzagMap]) -> ZigzagModuleId {
409    let mut hash = Sha256::new();
410    hash.update(b"holos-zigzag-module-v1");
411    hash.update(modulus.to_be_bytes());
412    hash.update((dimensions.len() as u64).to_be_bytes());
413    for dimension in dimensions {
414        hash.update((*dimension as u64).to_be_bytes());
415    }
416    for map in maps {
417        hash.update([match map.direction {
418            ZigzagDirection::Forward => 1,
419            ZigzagDirection::Backward => 2,
420        }]);
421        hash.update((map.columns.len() as u64).to_be_bytes());
422        for column in &map.columns {
423            hash.update((column.len() as u64).to_be_bytes());
424            for term in column {
425                hash.update((term.target as u64).to_be_bytes());
426                hash.update(term.coefficient.to_be_bytes());
427            }
428        }
429    }
430    ZigzagModuleId(hash.finalize().into())
431}
432
433fn interval_id(module: ZigzagModuleId, start: usize, end: usize) -> ZigzagIntervalId {
434    let mut hash = Sha256::new();
435    hash.update(b"holos-zigzag-interval-v1");
436    hash.update(module.as_bytes());
437    hash.update((start as u64).to_be_bytes());
438    hash.update((end as u64).to_be_bytes());
439    ZigzagIntervalId(hash.finalize().into())
440}
441
442fn write_hex(formatter: &mut fmt::Formatter<'_>, bytes: &[u8; 32]) -> fmt::Result {
443    for byte in bytes {
444        write!(formatter, "{byte:02x}")?;
445    }
446    Ok(())
447}
448
449#[cfg(test)]
450mod tests;