Skip to main content

sim_lib_numbers_stats/
hmm_fit.rs

1//! Deterministically initialized, bounded hidden Markov model fitting.
2
3use super::hmm_baum_welch::baum_welch_step;
4use super::hmm_inference::forward_backward;
5use super::hmm_model::{HiddenMarkovModel, HmmError};
6
7/// Stable numeric hidden-state identifier produced by [`fit_hmm`].
8pub type StateId = usize;
9
10/// One homogeneous observation sequence accepted by HMM fitting.
11#[derive(Clone, Debug, PartialEq)]
12pub enum Sequence {
13    /// Categorical symbols indexed from zero.
14    Discrete(Vec<usize>),
15    /// Finite scalar observations.
16    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/// Hidden-state and emission family requested from [`fit_hmm`].
29#[derive(Clone, Debug, PartialEq)]
30pub enum HmmSpec {
31    /// A finite model with categorical emissions.
32    Discrete {
33        /// Number of hidden states.
34        states: usize,
35        /// Number of observation symbols.
36        symbols: usize,
37        /// Pseudo-count added to fitted initial, transition, and emission rows.
38        additive_smoothing: f64,
39    },
40    /// A finite model with scalar Gaussian emissions.
41    Gaussian {
42        /// Number of hidden states.
43        states: usize,
44        /// Pseudo-count added to fitted initial and transition rows.
45        additive_smoothing: f64,
46        /// Hard lower bound for fitted variances.
47        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/// Deterministic initialization, convergence, and work policy for Baum-Welch.
101#[derive(Clone, Copy, Debug, PartialEq)]
102pub struct HmmFitControl {
103    /// Caller-owned deterministic initialization seed.
104    pub seed: u64,
105    /// Hard maximum number of accepted Baum-Welch updates.
106    pub max_iterations: usize,
107    /// Relative log-likelihood convergence tolerance.
108    pub tolerance: f64,
109    /// Hard maximum charged state-transition work.
110    pub max_work: u64,
111    /// Probability floor used while normalizing fitted rows.
112    pub probability_floor: f64,
113}
114
115impl HmmFitControl {
116    /// Builds checked fitting control.
117    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/// Why bounded Baum-Welch stopped.
165#[derive(Clone, Copy, Debug, PartialEq, Eq)]
166pub enum HmmTermination {
167    /// Relative log-likelihood improvement met the tolerance.
168    Converged,
169    /// The configured iteration count was exhausted.
170    IterationLimit,
171    /// The next complete expectation/maximization update exceeded `max_work`.
172    WorkLimit,
173    /// A candidate update reduced likelihood beyond numerical tolerance and
174    /// was rejected, leaving the last non-decreasing model in the report.
175    LikelihoodDecrease,
176}
177
178/// Convergence, likelihood, repair, seed, and termination evidence.
179#[derive(Clone, Debug, PartialEq)]
180pub struct HmmFitEvidence {
181    /// Initial model log likelihood before any update.
182    pub initial_log_likelihood: f64,
183    /// Final accepted model log likelihood.
184    pub log_likelihood: f64,
185    /// Initial value plus every accepted likelihood, in order.
186    pub likelihood_history: Vec<f64>,
187    /// Number of accepted Baum-Welch updates.
188    pub iterations: usize,
189    /// Whether convergence tolerance caused termination.
190    pub converged: bool,
191    /// Count of probability or variance floor repairs.
192    pub numerical_repairs: u64,
193    /// Caller-supplied initialization seed.
194    pub seed: u64,
195    /// Charged state-transition work, never greater than the control bound.
196    pub work: u64,
197    /// Concrete reason fitting stopped.
198    pub termination: HmmTermination,
199}
200
201/// A fitted hidden-state model together with complete termination evidence.
202#[derive(Clone, Debug, PartialEq)]
203pub struct HmmFitReport<M> {
204    /// Last accepted inspectable model.
205    pub model: M,
206    /// Fitting evidence and bounds.
207    pub evidence: HmmFitEvidence,
208}
209
210/// Fits a discrete- or continuous-emission HMM with bounded Baum-Welch.
211pub 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}