1use super::hmm_baum_welch::baum_welch_step;
4use super::hmm_inference::forward_backward;
5use super::hmm_model::{HiddenMarkovModel, HmmError};
6
7pub type StateId = usize;
9
10#[derive(Clone, Debug, PartialEq)]
12pub enum Sequence {
13 Discrete(Vec<usize>),
15 Continuous(Vec<f64>),
17}
18
19impl Sequence {
20 fn len(&self) -> usize {
21 match self {
22 Self::Discrete(values) => values.len(),
23 Self::Continuous(values) => values.len(),
24 }
25 }
26}
27
28#[derive(Clone, Debug, PartialEq)]
30pub enum HmmSpec {
31 Discrete {
33 states: usize,
35 symbols: usize,
37 additive_smoothing: f64,
39 },
40 Gaussian {
42 states: usize,
44 additive_smoothing: f64,
46 variance_floor: f64,
48 },
49}
50
51impl HmmSpec {
52 pub(crate) fn states(&self) -> usize {
53 match self {
54 Self::Discrete { states, .. } | Self::Gaussian { states, .. } => *states,
55 }
56 }
57
58 pub(crate) fn smoothing(&self) -> f64 {
59 match self {
60 Self::Discrete {
61 additive_smoothing, ..
62 }
63 | Self::Gaussian {
64 additive_smoothing, ..
65 } => *additive_smoothing,
66 }
67 }
68
69 fn validate(&self) -> Result<(), HmmError> {
70 if self.states() == 0 {
71 return Err(HmmError::InvalidFitControl {
72 field: "spec.states",
73 reason: "must be greater than zero",
74 });
75 }
76 if !self.smoothing().is_finite() || self.smoothing() <= 0.0 {
77 return Err(HmmError::InvalidFitControl {
78 field: "spec.additive_smoothing",
79 reason: "must be finite and greater than zero",
80 });
81 }
82 match self {
83 Self::Discrete { symbols: 0, .. } => Err(HmmError::InvalidFitControl {
84 field: "spec.symbols",
85 reason: "must be greater than zero",
86 }),
87 Self::Gaussian { variance_floor, .. }
88 if !variance_floor.is_finite() || *variance_floor <= 0.0 =>
89 {
90 Err(HmmError::InvalidFitControl {
91 field: "spec.variance_floor",
92 reason: "must be finite and greater than zero",
93 })
94 }
95 _ => Ok(()),
96 }
97 }
98}
99
100#[derive(Clone, Copy, Debug, PartialEq)]
102pub struct HmmFitControl {
103 pub seed: u64,
105 pub max_iterations: usize,
107 pub tolerance: f64,
109 pub max_work: u64,
111 pub probability_floor: f64,
113}
114
115impl HmmFitControl {
116 pub fn new(
118 seed: u64,
119 max_iterations: usize,
120 tolerance: f64,
121 max_work: u64,
122 probability_floor: f64,
123 ) -> Result<Self, HmmError> {
124 let control = Self {
125 seed,
126 max_iterations,
127 tolerance,
128 max_work,
129 probability_floor,
130 };
131 control.validate()?;
132 Ok(control)
133 }
134
135 fn validate(&self) -> Result<(), HmmError> {
136 for (field, valid, reason) in [
137 (
138 "max_iterations",
139 self.max_iterations > 0,
140 "must be greater than zero",
141 ),
142 ("max_work", self.max_work > 0, "must be greater than zero"),
143 (
144 "tolerance",
145 self.tolerance.is_finite() && self.tolerance >= 0.0,
146 "must be finite and nonnegative",
147 ),
148 (
149 "probability_floor",
150 self.probability_floor.is_finite()
151 && self.probability_floor > 0.0
152 && self.probability_floor < 1.0,
153 "must be finite and in the open interval (0, 1)",
154 ),
155 ] {
156 if !valid {
157 return Err(HmmError::InvalidFitControl { field, reason });
158 }
159 }
160 Ok(())
161 }
162}
163
164#[derive(Clone, Copy, Debug, PartialEq, Eq)]
166pub enum HmmTermination {
167 Converged,
169 IterationLimit,
171 WorkLimit,
173 LikelihoodDecrease,
176}
177
178#[derive(Clone, Debug, PartialEq)]
180pub struct HmmFitEvidence {
181 pub initial_log_likelihood: f64,
183 pub log_likelihood: f64,
185 pub likelihood_history: Vec<f64>,
187 pub iterations: usize,
189 pub converged: bool,
191 pub numerical_repairs: u64,
193 pub seed: u64,
195 pub work: u64,
197 pub termination: HmmTermination,
199}
200
201#[derive(Clone, Debug, PartialEq)]
203pub struct HmmFitReport<M> {
204 pub model: M,
206 pub evidence: HmmFitEvidence,
208}
209
210pub fn fit_hmm(
212 data: &[Sequence],
213 spec: HmmSpec,
214 control: HmmFitControl,
215) -> Result<HmmFitReport<HiddenMarkovModel<StateId>>, HmmError> {
216 spec.validate()?;
217 control.validate()?;
218 validate_data(data, &spec)?;
219 let unit_work = inference_work(data, spec.states())?;
220 if unit_work > control.max_work {
221 return Err(HmmError::InvalidFitControl {
222 field: "max_work",
223 reason: "must admit the initial likelihood sweep",
224 });
225 }
226 let mut model = initialize_model(data, &spec, control.seed)?;
227 let initial_log_likelihood = score_data(&model, data)?;
228 let mut history = vec![initial_log_likelihood];
229 let mut work = unit_work;
230 let mut iterations = 0;
231 let mut numerical_repairs = 0_u64;
232
233 let termination = loop {
234 if iterations == control.max_iterations {
235 break HmmTermination::IterationLimit;
236 }
237 let update_work = unit_work.checked_mul(2).ok_or(HmmError::WorkOverflow)?;
238 if work
239 .checked_add(update_work)
240 .is_none_or(|next| next > control.max_work)
241 {
242 break HmmTermination::WorkLimit;
243 }
244 let (candidate, repairs) = baum_welch_step(&model, data, &spec, control.probability_floor)?;
245 let likelihood = score_data(&candidate, data)?;
246 work += update_work;
247 numerical_repairs = numerical_repairs.saturating_add(repairs);
248 let previous = *history.last().unwrap_or(&initial_log_likelihood);
249 let scale = previous.abs().max(1.0);
250 if likelihood + control.tolerance * scale < previous {
251 break HmmTermination::LikelihoodDecrease;
252 }
253 model = candidate;
254 history.push(likelihood);
255 iterations += 1;
256 if (likelihood - previous).abs() <= control.tolerance * scale {
257 break HmmTermination::Converged;
258 }
259 };
260
261 let log_likelihood = *history.last().unwrap_or(&initial_log_likelihood);
262 Ok(HmmFitReport {
263 model,
264 evidence: HmmFitEvidence {
265 initial_log_likelihood,
266 log_likelihood,
267 likelihood_history: history,
268 iterations,
269 converged: termination == HmmTermination::Converged,
270 numerical_repairs,
271 seed: control.seed,
272 work,
273 termination,
274 },
275 })
276}
277
278fn validate_data(data: &[Sequence], spec: &HmmSpec) -> Result<(), HmmError> {
279 if data.is_empty() {
280 return Err(HmmError::EmptyInput);
281 }
282 for (index, sequence) in data.iter().enumerate() {
283 if sequence.len() == 0 {
284 return Err(HmmError::EmptySequence { index });
285 }
286 match (sequence, spec) {
287 (Sequence::Discrete(values), HmmSpec::Discrete { symbols, .. }) => {
288 if let Some(&symbol) = values.iter().find(|&&symbol| symbol >= *symbols) {
289 return Err(HmmError::UnknownSymbol {
290 symbol,
291 symbol_count: *symbols,
292 });
293 }
294 }
295 (Sequence::Continuous(values), HmmSpec::Gaussian { .. }) => {
296 if let Some(&value) = values.iter().find(|value| !value.is_finite()) {
297 return Err(HmmError::NonFiniteObservation { value });
298 }
299 }
300 _ => return Err(HmmError::MixedSequenceKinds),
301 }
302 }
303 Ok(())
304}
305
306fn inference_work(data: &[Sequence], states: usize) -> Result<u64, HmmError> {
307 let observations = data.iter().try_fold(0_u64, |sum, sequence| {
308 sum.checked_add(sequence.len() as u64)
309 .ok_or(HmmError::WorkOverflow)
310 })?;
311 observations
312 .checked_mul(states as u64)
313 .and_then(|value| value.checked_mul(states as u64))
314 .ok_or(HmmError::WorkOverflow)
315}
316
317fn score_data(model: &HiddenMarkovModel<StateId>, data: &[Sequence]) -> Result<f64, HmmError> {
318 data.iter().try_fold(0.0, |sum, sequence| {
319 let likelihood = match sequence {
320 Sequence::Discrete(values) => forward_backward(model, values)?.evidence.log_likelihood,
321 Sequence::Continuous(values) => {
322 forward_backward(model, values)?.evidence.log_likelihood
323 }
324 };
325 Ok(sum + likelihood)
326 })
327}
328
329fn initialize_model(
330 data: &[Sequence],
331 spec: &HmmSpec,
332 seed: u64,
333) -> Result<HiddenMarkovModel<StateId>, HmmError> {
334 let states = spec.states();
335 let state_ids = (0..states).collect::<Vec<_>>();
336 let mut random = SplitMix64::new(seed);
337 let initial = random_distribution(states, &mut random);
338 let transitions = (0..states)
339 .map(|_| random_distribution(states, &mut random))
340 .collect::<Vec<_>>();
341 match spec {
342 HmmSpec::Discrete { symbols, .. } => {
343 let emissions = (0..states)
344 .map(|_| random_distribution(*symbols, &mut random))
345 .collect();
346 HiddenMarkovModel::discrete(state_ids, initial, transitions, emissions)
347 }
348 HmmSpec::Gaussian { variance_floor, .. } => {
349 let values = data
350 .iter()
351 .flat_map(|sequence| match sequence {
352 Sequence::Continuous(values) => values.as_slice(),
353 Sequence::Discrete(_) => &[],
354 })
355 .copied()
356 .collect::<Vec<_>>();
357 let global_mean = values.iter().sum::<f64>() / values.len() as f64;
358 let global_variance = values
359 .iter()
360 .map(|value| (value - global_mean).powi(2))
361 .sum::<f64>()
362 / values.len() as f64;
363 let variance = global_variance.max(*variance_floor);
364 let means = (0..states)
365 .map(|_| values[random.index(values.len())])
366 .collect();
367 HiddenMarkovModel::gaussian(
368 state_ids,
369 initial,
370 transitions,
371 means,
372 vec![variance; states],
373 *variance_floor,
374 )
375 }
376 }
377}
378
379fn random_distribution(length: usize, random: &mut SplitMix64) -> Vec<f64> {
380 let mut values = (0..length)
381 .map(|_| 0.5 + random.unit_interval())
382 .collect::<Vec<_>>();
383 let sum = values.iter().sum::<f64>();
384 for value in &mut values {
385 *value /= sum;
386 }
387 values
388}
389
390struct SplitMix64 {
391 state: u64,
392}
393
394impl SplitMix64 {
395 fn new(seed: u64) -> Self {
396 Self { state: seed }
397 }
398
399 fn next(&mut self) -> u64 {
400 self.state = self.state.wrapping_add(0x9e3779b97f4a7c15);
401 let mut value = self.state;
402 value = (value ^ (value >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
403 value = (value ^ (value >> 27)).wrapping_mul(0x94d049bb133111eb);
404 value ^ (value >> 31)
405 }
406
407 fn unit_interval(&mut self) -> f64 {
408 (self.next() >> 11) as f64 * (1.0 / ((1_u64 << 53) as f64))
409 }
410
411 fn index(&mut self, length: usize) -> usize {
412 (self.next() % length as u64) as usize
413 }
414}