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