Skip to main content

sim_lib_numbers_stats/
markov.rs

1//! Finite, inspectable first-order Markov transition estimation.
2
3use super::transition::FiniteTransitionMatrix;
4use std::collections::{BTreeMap, BTreeSet};
5use std::error::Error;
6use std::fmt;
7
8/// Stable provenance attached to every fitted transition model.
9#[derive(Clone, Debug, PartialEq, Eq)]
10pub struct CorpusProvenance {
11    /// Stable corpus identifier.
12    pub id: String,
13    /// Human-readable origin, generator, or public source.
14    pub source: String,
15    /// SPDX license identifier or an equally precise public-domain declaration.
16    pub license: String,
17    /// Content hash of the exact corpus bytes.
18    pub content_hash: String,
19}
20
21impl CorpusProvenance {
22    /// Builds checked provenance from a previously computed content hash.
23    pub fn new(
24        id: impl Into<String>,
25        source: impl Into<String>,
26        license: impl Into<String>,
27        content_hash: impl Into<String>,
28    ) -> Result<Self, MarkovError> {
29        let provenance = Self {
30            id: id.into(),
31            source: source.into(),
32            license: license.into(),
33            content_hash: content_hash.into(),
34        };
35        provenance.validate()?;
36        Ok(provenance)
37    }
38
39    /// Builds provenance and computes an FNV-1a hash over the exact corpus bytes.
40    pub fn from_bytes(
41        id: impl Into<String>,
42        source: impl Into<String>,
43        license: impl Into<String>,
44        bytes: &[u8],
45    ) -> Result<Self, MarkovError> {
46        Self::new(id, source, license, fnv1a64(bytes))
47    }
48
49    fn validate(&self) -> Result<(), MarkovError> {
50        for (field, value) in [
51            ("corpus.id", self.id.as_str()),
52            ("corpus.source", self.source.as_str()),
53            ("corpus.license", self.license.as_str()),
54            ("corpus.content_hash", self.content_hash.as_str()),
55        ] {
56            if value.trim().is_empty() {
57                return Err(MarkovError::InvalidPolicy {
58                    field,
59                    reason: "must not be empty",
60                });
61            }
62        }
63        Ok(())
64    }
65}
66
67/// Explicit fitting and evaluation policy for a finite first-order model.
68#[derive(Clone, Debug, PartialEq)]
69pub struct MarkovPolicy {
70    /// Additive pseudo-count applied to every transition in the finite state set.
71    pub additive_smoothing: f64,
72    /// Number of trailing sequences reserved for held-out evaluation.
73    pub held_out_sequences: usize,
74    /// Identity, origin, license, and exact content hash of the corpus.
75    pub corpus: CorpusProvenance,
76}
77
78impl MarkovPolicy {
79    /// Builds a policy and rejects absent smoothing or incomplete provenance.
80    pub fn new(
81        additive_smoothing: f64,
82        held_out_sequences: usize,
83        corpus: CorpusProvenance,
84    ) -> Result<Self, MarkovError> {
85        let policy = Self {
86            additive_smoothing,
87            held_out_sequences,
88            corpus,
89        };
90        policy.validate()?;
91        Ok(policy)
92    }
93
94    fn validate(&self) -> Result<(), MarkovError> {
95        if !self.additive_smoothing.is_finite() || self.additive_smoothing <= 0.0 {
96            return Err(MarkovError::InvalidPolicy {
97                field: "additive_smoothing",
98                reason: "must be finite and greater than zero",
99            });
100        }
101        self.corpus.validate()
102    }
103}
104
105/// Aggregate likelihood evidence for a collection of state sequences.
106#[derive(Clone, Copy, Debug, PartialEq)]
107pub struct TransitionScore {
108    /// Number of adjacent transitions scored.
109    pub transitions: u64,
110    /// Sum of natural logarithms of smoothed transition probabilities.
111    pub log_likelihood: f64,
112    /// Mean negative natural-log likelihood per transition.
113    pub mean_negative_log_likelihood: f64,
114    /// Exponential of the mean negative log likelihood.
115    pub perplexity: f64,
116}
117
118/// A fitted value together with training and held-out evidence.
119#[derive(Clone, Debug, PartialEq)]
120pub struct ModelReport<M> {
121    /// Inspectable fitted model.
122    pub model: M,
123    /// Number of sequences used to estimate transition counts.
124    pub training_sequences: usize,
125    /// Number of trailing sequences excluded from estimation.
126    pub held_out_sequences: usize,
127    /// Likelihood of the exact training partition under the fitted model.
128    pub training_score: TransitionScore,
129    /// Likelihood of the held-out partition, when one was requested.
130    pub held_out_score: Option<TransitionScore>,
131}
132
133/// A finite first-order Markov model retaining exact counts and fitting policy.
134#[derive(Clone, Debug, PartialEq)]
135pub struct MarkovModel<S> {
136    states: Vec<S>,
137    transition_counts: BTreeMap<(S, S), u64>,
138    outgoing_counts: BTreeMap<S, u64>,
139    policy: MarkovPolicy,
140}
141
142impl<S: Ord + Clone> MarkovModel<S> {
143    /// Returns the sorted finite state vocabulary.
144    pub fn states(&self) -> &[S] {
145        &self.states
146    }
147
148    /// Returns the policy and corpus provenance used during fitting.
149    pub fn policy(&self) -> &MarkovPolicy {
150        &self.policy
151    }
152
153    /// Projects the fitted counts and smoothing policy into the shared finite
154    /// transition representation used by hidden-state inference.
155    pub fn transition_matrix(&self) -> FiniteTransitionMatrix<S> {
156        let probabilities = self
157            .states
158            .iter()
159            .map(|from| {
160                self.states
161                    .iter()
162                    .map(|to| {
163                        let count = self
164                            .transition_counts
165                            .get(&(from.clone(), to.clone()))
166                            .copied()
167                            .unwrap_or(0) as f64;
168                        let outgoing = self.outgoing_counts.get(from).copied().unwrap_or(0) as f64;
169                        let smoothing = self.policy.additive_smoothing;
170                        (count + smoothing) / (outgoing + smoothing * self.states.len() as f64)
171                    })
172                    .collect()
173            })
174            .collect();
175        FiniteTransitionMatrix::from_normalized(self.states.clone(), probabilities)
176    }
177
178    /// Returns the exact observed count for one transition.
179    pub fn transition_count(&self, from: &S, to: &S) -> Result<u64, MarkovError> {
180        self.require_state(from, 0, 0)?;
181        self.require_state(to, 0, 1)?;
182        Ok(self
183            .transition_counts
184            .get(&(from.clone(), to.clone()))
185            .copied()
186            .unwrap_or(0))
187    }
188
189    /// Returns the additively smoothed probability of one transition.
190    pub fn transition_probability(&self, from: &S, to: &S) -> Result<f64, MarkovError> {
191        let count = self.transition_count(from, to)? as f64;
192        let outgoing = self.outgoing_counts.get(from).copied().unwrap_or(0) as f64;
193        let smoothing = self.policy.additive_smoothing;
194        Ok((count + smoothing) / (outgoing + smoothing * self.states.len() as f64))
195    }
196
197    /// Scores sequences without changing the fitted model.
198    pub fn score(&self, sequences: &[Vec<S>]) -> Result<TransitionScore, MarkovError> {
199        score_sequences(self, sequences, "evaluation")
200    }
201
202    /// Serializes policy, provenance, states, and exact counts deterministically.
203    ///
204    /// The caller supplies a stable, domain-owned state label. Labels are
205    /// hex-encoded and indexed in sorted state order, so punctuation and
206    /// whitespace cannot make the representation ambiguous.
207    pub fn to_stable_text(
208        &self,
209        mut state_label: impl FnMut(&S) -> String,
210    ) -> Result<String, MarkovError> {
211        let labels = self.states.iter().map(&mut state_label).collect::<Vec<_>>();
212        let unique = labels.iter().collect::<BTreeSet<_>>();
213        if unique.len() != labels.len() {
214            return Err(MarkovError::DuplicateStateLabel);
215        }
216        let indices = self
217            .states
218            .iter()
219            .cloned()
220            .enumerate()
221            .map(|(index, state)| (state, index))
222            .collect::<BTreeMap<_, _>>();
223        let mut text = String::from("SIM-MARKOV-1\n");
224        text.push_str(&format!(
225            "additive-smoothing-bits={:016x}\n",
226            self.policy.additive_smoothing.to_bits()
227        ));
228        text.push_str(&format!(
229            "held-out-sequences={}\n",
230            self.policy.held_out_sequences
231        ));
232        for (name, value) in [
233            ("corpus-id", self.policy.corpus.id.as_str()),
234            ("corpus-source", self.policy.corpus.source.as_str()),
235            ("corpus-license", self.policy.corpus.license.as_str()),
236            ("corpus-hash", self.policy.corpus.content_hash.as_str()),
237        ] {
238            text.push_str(name);
239            text.push('=');
240            text.push_str(&hex(value.as_bytes()));
241            text.push('\n');
242        }
243        text.push_str(&format!("states={}\n", labels.len()));
244        for (index, label) in labels.iter().enumerate() {
245            text.push_str(&format!("state={index}:{}\n", hex(label.as_bytes())));
246        }
247        for ((from, to), count) in &self.transition_counts {
248            text.push_str(&format!(
249                "transition={}:{}:{count}\n",
250                indices[from], indices[to]
251            ));
252        }
253        Ok(text)
254    }
255
256    fn require_state(
257        &self,
258        state: &S,
259        sequence: usize,
260        position: usize,
261    ) -> Result<(), MarkovError> {
262        if self.states.binary_search(state).is_err() {
263            return Err(MarkovError::UnknownState { sequence, position });
264        }
265        Ok(())
266    }
267}
268
269/// Fits an inspectable finite first-order model and reports held-out evidence.
270///
271/// The last `policy.held_out_sequences` are excluded from both the finite
272/// vocabulary and count estimation. A held-out state absent from training
273/// therefore fails closed instead of leaking evaluation data into the model.
274pub fn fit_markov<S: Ord + Clone>(
275    sequences: &[Vec<S>],
276    policy: MarkovPolicy,
277) -> Result<ModelReport<MarkovModel<S>>, MarkovError> {
278    policy.validate()?;
279    if sequences.is_empty() {
280        return Err(MarkovError::EmptyCorpus);
281    }
282    if policy.held_out_sequences >= sequences.len() {
283        return Err(MarkovError::InvalidHoldout {
284            sequences: sequences.len(),
285            held_out: policy.held_out_sequences,
286        });
287    }
288    for (index, sequence) in sequences.iter().enumerate() {
289        if sequence.is_empty() {
290            return Err(MarkovError::EmptySequence { index });
291        }
292    }
293
294    let split = sequences.len() - policy.held_out_sequences;
295    require_transitions(&sequences[..split], "training")?;
296    if split < sequences.len() {
297        require_transitions(&sequences[split..], "held-out")?;
298    }
299    let states = sequences[..split]
300        .iter()
301        .flat_map(|sequence| sequence.iter().cloned())
302        .collect::<BTreeSet<_>>()
303        .into_iter()
304        .collect::<Vec<_>>();
305    let mut transition_counts = BTreeMap::new();
306    let mut outgoing_counts = BTreeMap::new();
307    for sequence in &sequences[..split] {
308        for pair in sequence.windows(2) {
309            increment(
310                transition_counts
311                    .entry((pair[0].clone(), pair[1].clone()))
312                    .or_insert(0),
313            )?;
314            increment(outgoing_counts.entry(pair[0].clone()).or_insert(0))?;
315        }
316    }
317    let model = MarkovModel {
318        states,
319        transition_counts,
320        outgoing_counts,
321        policy,
322    };
323    let training_score = score_sequences(&model, &sequences[..split], "training")?;
324    let held_out_score = if split < sequences.len() {
325        Some(score_sequences(&model, &sequences[split..], "held-out")?)
326    } else {
327        None
328    };
329    Ok(ModelReport {
330        model,
331        training_sequences: split,
332        held_out_sequences: sequences.len() - split,
333        training_score,
334        held_out_score,
335    })
336}
337
338/// Computes a stable FNV-1a digest for small transparent fixture corpora.
339pub fn fnv1a64(bytes: &[u8]) -> String {
340    let mut hash = 0xcbf29ce484222325_u64;
341    for byte in bytes {
342        hash ^= u64::from(*byte);
343        hash = hash.wrapping_mul(0x100000001b3);
344    }
345    format!("fnv1a64:{hash:016x}")
346}
347
348/// Failure while validating, fitting, scoring, or serializing a Markov model.
349#[derive(Clone, Debug, PartialEq, Eq)]
350pub enum MarkovError {
351    /// The corpus contained no sequences.
352    EmptyCorpus,
353    /// One sequence contained no state.
354    EmptySequence {
355        /// Zero-based sequence position.
356        index: usize,
357    },
358    /// The requested partition contained no adjacent state pair.
359    NoTransitions {
360        /// Stable partition name.
361        partition: &'static str,
362    },
363    /// The holdout would leave no training sequence.
364    InvalidHoldout {
365        /// Total sequence count.
366        sequences: usize,
367        /// Requested held-out count.
368        held_out: usize,
369    },
370    /// Policy or provenance was incomplete or numerically invalid.
371    InvalidPolicy {
372        /// Invalid field.
373        field: &'static str,
374        /// Concrete requirement.
375        reason: &'static str,
376    },
377    /// A scored state was outside the fitted finite vocabulary.
378    UnknownState {
379        /// Zero-based sequence position.
380        sequence: usize,
381        /// Zero-based state position.
382        position: usize,
383    },
384    /// A transition count exceeded `u64`.
385    CountOverflow,
386    /// Two distinct states were given the same stable serialization label.
387    DuplicateStateLabel,
388}
389
390impl fmt::Display for MarkovError {
391    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
392        match self {
393            Self::EmptyCorpus => write!(formatter, "Markov fitting requires at least one sequence"),
394            Self::EmptySequence { index } => {
395                write!(formatter, "Markov sequence {index} contains no state")
396            }
397            Self::NoTransitions { partition } => {
398                write!(
399                    formatter,
400                    "Markov {partition} partition contains no transition"
401                )
402            }
403            Self::InvalidHoldout {
404                sequences,
405                held_out,
406            } => write!(
407                formatter,
408                "Markov holdout {held_out} must be smaller than sequence count {sequences}"
409            ),
410            Self::InvalidPolicy { field, reason } => {
411                write!(formatter, "invalid Markov policy {field}: {reason}")
412            }
413            Self::UnknownState { sequence, position } => write!(
414                formatter,
415                "Markov sequence {sequence} state {position} is outside the finite vocabulary"
416            ),
417            Self::CountOverflow => write!(formatter, "Markov transition count overflow"),
418            Self::DuplicateStateLabel => {
419                write!(formatter, "Markov stable state labels must be unique")
420            }
421        }
422    }
423}
424
425impl Error for MarkovError {}
426
427fn score_sequences<S: Ord + Clone>(
428    model: &MarkovModel<S>,
429    sequences: &[Vec<S>],
430    partition: &'static str,
431) -> Result<TransitionScore, MarkovError> {
432    let mut transitions = 0_u64;
433    let mut log_likelihood = 0.0;
434    for (sequence_index, sequence) in sequences.iter().enumerate() {
435        for (position, state) in sequence.iter().enumerate() {
436            model.require_state(state, sequence_index, position)?;
437        }
438        for pair in sequence.windows(2) {
439            log_likelihood += model.transition_probability(&pair[0], &pair[1])?.ln();
440            increment(&mut transitions)?;
441        }
442    }
443    if transitions == 0 {
444        return Err(MarkovError::NoTransitions { partition });
445    }
446    let mean_negative_log_likelihood = -log_likelihood / transitions as f64;
447    Ok(TransitionScore {
448        transitions,
449        log_likelihood,
450        mean_negative_log_likelihood,
451        perplexity: mean_negative_log_likelihood.exp(),
452    })
453}
454
455fn require_transitions<S>(
456    sequences: &[Vec<S>],
457    partition: &'static str,
458) -> Result<(), MarkovError> {
459    if sequences.iter().all(|sequence| sequence.len() < 2) {
460        return Err(MarkovError::NoTransitions { partition });
461    }
462    Ok(())
463}
464
465fn increment(value: &mut u64) -> Result<(), MarkovError> {
466    *value = value.checked_add(1).ok_or(MarkovError::CountOverflow)?;
467    Ok(())
468}
469
470fn hex(bytes: &[u8]) -> String {
471    const DIGITS: &[u8; 16] = b"0123456789abcdef";
472    let mut encoded = String::with_capacity(bytes.len() * 2);
473    for byte in bytes {
474        encoded.push(DIGITS[(byte >> 4) as usize] as char);
475        encoded.push(DIGITS[(byte & 0x0f) as usize] as char);
476    }
477    encoded
478}