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