1use super::transition::FiniteTransitionMatrix;
4use std::collections::{BTreeMap, BTreeSet};
5use std::error::Error;
6use std::fmt;
7
8#[derive(Clone, Debug, PartialEq, Eq)]
10pub struct CorpusProvenance {
11 pub id: String,
13 pub source: String,
15 pub license: String,
17 pub content_hash: String,
19}
20
21impl CorpusProvenance {
22 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 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#[derive(Clone, Debug, PartialEq)]
69pub struct MarkovPolicy {
70 pub additive_smoothing: f64,
72 pub held_out_sequences: usize,
74 pub corpus: CorpusProvenance,
76}
77
78impl MarkovPolicy {
79 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#[derive(Clone, Copy, Debug, PartialEq)]
107pub struct TransitionScore {
108 pub transitions: u64,
110 pub log_likelihood: f64,
112 pub mean_negative_log_likelihood: f64,
114 pub perplexity: f64,
116}
117
118#[derive(Clone, Debug, PartialEq)]
120pub struct ModelReport<M> {
121 pub model: M,
123 pub training_sequences: usize,
125 pub held_out_sequences: usize,
127 pub training_score: TransitionScore,
129 pub held_out_score: Option<TransitionScore>,
131}
132
133#[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 pub fn states(&self) -> &[S] {
145 &self.states
146 }
147
148 pub fn policy(&self) -> &MarkovPolicy {
150 &self.policy
151 }
152
153 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 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 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 pub fn score(&self, sequences: &[Vec<S>]) -> Result<TransitionScore, MarkovError> {
199 score_sequences(self, sequences, "evaluation")
200 }
201
202 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
269pub 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
338pub 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#[derive(Clone, Debug, PartialEq, Eq)]
350pub enum MarkovError {
351 EmptyCorpus,
353 EmptySequence {
355 index: usize,
357 },
358 NoTransitions {
360 partition: &'static str,
362 },
363 InvalidHoldout {
365 sequences: usize,
367 held_out: usize,
369 },
370 InvalidPolicy {
372 field: &'static str,
374 reason: &'static str,
376 },
377 UnknownState {
379 sequence: usize,
381 position: usize,
383 },
384 CountOverflow,
386 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}