Skip to main content

sim_lib_numbers_stats/
hmm_model.rs

1//! Finite hidden-state model and emission representations.
2
3use super::transition::{FiniteTransitionMatrix, TransitionError, validate_distribution};
4use std::{error::Error, f64::consts::TAU, fmt};
5
6/// Discrete or scalar Gaussian emissions for a finite hidden-state model.
7#[derive(Clone, Debug, PartialEq)]
8pub enum EmissionModel {
9    /// A row-stochastic categorical distribution for every hidden state.
10    Discrete {
11        /// State-major probability rows over symbols `0..symbol_count`.
12        probabilities: Vec<Vec<f64>>,
13    },
14    /// One univariate Gaussian distribution for every hidden state.
15    Gaussian {
16        /// Mean for each hidden state.
17        means: Vec<f64>,
18        /// Strictly positive variance for each hidden state.
19        variances: Vec<f64>,
20        /// Fitting floor retained with the model as numerical policy.
21        variance_floor: f64,
22    },
23}
24
25impl EmissionModel {
26    /// Returns the number of categorical symbols, or `None` for Gaussian emissions.
27    pub fn symbol_count(&self) -> Option<usize> {
28        match self {
29            Self::Discrete { probabilities } => probabilities.first().map(Vec::len),
30            Self::Gaussian { .. } => None,
31        }
32    }
33
34    /// Returns the state-major categorical rows, when this is a discrete model.
35    pub fn discrete_probabilities(&self) -> Option<&[Vec<f64>]> {
36        match self {
37            Self::Discrete { probabilities } => Some(probabilities),
38            Self::Gaussian { .. } => None,
39        }
40    }
41
42    /// Returns Gaussian means, variances, and variance floor when applicable.
43    pub fn gaussian_parameters(&self) -> Option<(&[f64], &[f64], f64)> {
44        match self {
45            Self::Gaussian {
46                means,
47                variances,
48                variance_floor,
49            } => Some((means, variances, *variance_floor)),
50            Self::Discrete { .. } => None,
51        }
52    }
53
54    fn validate(&self, states: usize) -> Result<(), HmmError> {
55        match self {
56            Self::Discrete { probabilities } => {
57                if probabilities.len() != states {
58                    return Err(HmmError::EmissionStateCount {
59                        expected: states,
60                        actual: probabilities.len(),
61                    });
62                }
63                let symbols = probabilities.first().map_or(0, Vec::len);
64                if symbols == 0 {
65                    return Err(HmmError::InvalidModel {
66                        field: "emission.symbols",
67                        reason: "must be greater than zero",
68                    });
69                }
70                for (state, row) in probabilities.iter().enumerate() {
71                    validate_distribution("emission", state, row, symbols)?;
72                }
73            }
74            Self::Gaussian {
75                means,
76                variances,
77                variance_floor,
78            } => {
79                if means.len() != states || variances.len() != states {
80                    return Err(HmmError::EmissionStateCount {
81                        expected: states,
82                        actual: means.len().min(variances.len()),
83                    });
84                }
85                if !variance_floor.is_finite() || *variance_floor <= 0.0 {
86                    return Err(HmmError::InvalidModel {
87                        field: "emission.variance_floor",
88                        reason: "must be finite and greater than zero",
89                    });
90                }
91                for (state, (&mean, &variance)) in means.iter().zip(variances).enumerate() {
92                    if !mean.is_finite() {
93                        return Err(HmmError::InvalidGaussian {
94                            state,
95                            field: "mean",
96                            value: mean,
97                        });
98                    }
99                    if !variance.is_finite() || variance < *variance_floor {
100                        return Err(HmmError::InvalidGaussian {
101                            state,
102                            field: "variance",
103                            value: variance,
104                        });
105                    }
106                }
107            }
108        }
109        Ok(())
110    }
111
112    pub(crate) fn log_discrete(&self, state: usize, symbol: usize) -> Result<f64, HmmError> {
113        let Self::Discrete { probabilities } = self else {
114            return Err(HmmError::EmissionKind {
115                expected: "discrete",
116                actual: "continuous",
117            });
118        };
119        let symbol_count = probabilities.first().map_or(0, Vec::len);
120        if symbol >= symbol_count {
121            return Err(HmmError::UnknownSymbol {
122                symbol,
123                symbol_count,
124            });
125        }
126        Ok(log_probability(probabilities[state][symbol]))
127    }
128
129    pub(crate) fn log_continuous(&self, state: usize, value: f64) -> Result<f64, HmmError> {
130        if !value.is_finite() {
131            return Err(HmmError::NonFiniteObservation { value });
132        }
133        let Self::Gaussian {
134            means, variances, ..
135        } = self
136        else {
137            return Err(HmmError::EmissionKind {
138                expected: "continuous",
139                actual: "discrete",
140            });
141        };
142        let difference = value - means[state];
143        Ok(-0.5 * (difference * difference / variances[state] + (TAU * variances[state]).ln()))
144    }
145}
146
147/// Observation accepted by generic HMM inference.
148pub trait HmmObservation: Copy {
149    /// Returns this observation's log likelihood in `state`.
150    fn emission_log_likelihood(
151        self,
152        emissions: &EmissionModel,
153        state: usize,
154    ) -> Result<f64, HmmError>;
155}
156
157impl HmmObservation for usize {
158    fn emission_log_likelihood(
159        self,
160        emissions: &EmissionModel,
161        state: usize,
162    ) -> Result<f64, HmmError> {
163        emissions.log_discrete(state, self)
164    }
165}
166
167impl HmmObservation for f64 {
168    fn emission_log_likelihood(
169        self,
170        emissions: &EmissionModel,
171        state: usize,
172    ) -> Result<f64, HmmError> {
173        emissions.log_continuous(state, self)
174    }
175}
176
177/// A finite hidden Markov model with inspectable transition and emission rows.
178#[derive(Clone, Debug, PartialEq)]
179pub struct HiddenMarkovModel<S> {
180    initial: Vec<f64>,
181    transitions: FiniteTransitionMatrix<S>,
182    emissions: EmissionModel,
183}
184
185impl<S: Eq + Clone> HiddenMarkovModel<S> {
186    /// Builds a model with categorical observations indexed from zero.
187    pub fn discrete(
188        states: Vec<S>,
189        initial: Vec<f64>,
190        transitions: Vec<Vec<f64>>,
191        emissions: Vec<Vec<f64>>,
192    ) -> Result<Self, HmmError> {
193        Self::from_transition_matrix(
194            initial,
195            FiniteTransitionMatrix::new(states, transitions)?,
196            EmissionModel::Discrete {
197                probabilities: emissions,
198            },
199        )
200    }
201
202    /// Builds a model with scalar Gaussian observations.
203    pub fn gaussian(
204        states: Vec<S>,
205        initial: Vec<f64>,
206        transitions: Vec<Vec<f64>>,
207        means: Vec<f64>,
208        variances: Vec<f64>,
209        variance_floor: f64,
210    ) -> Result<Self, HmmError> {
211        Self::from_transition_matrix(
212            initial,
213            FiniteTransitionMatrix::new(states, transitions)?,
214            EmissionModel::Gaussian {
215                means,
216                variances,
217                variance_floor,
218            },
219        )
220    }
221
222    /// Builds a model from the same finite transition representation exposed
223    /// by [`crate::MarkovModel::transition_matrix`].
224    pub fn from_transition_matrix(
225        initial: Vec<f64>,
226        transitions: FiniteTransitionMatrix<S>,
227        emissions: EmissionModel,
228    ) -> Result<Self, HmmError> {
229        validate_distribution("initial", 0, &initial, transitions.len())?;
230        emissions.validate(transitions.len())?;
231        Ok(Self {
232            initial,
233            transitions,
234            emissions,
235        })
236    }
237
238    /// Returns the ordered hidden-state vocabulary.
239    pub fn states(&self) -> &[S] {
240        self.transitions.states()
241    }
242
243    /// Returns the normalized initial-state probabilities.
244    pub fn initial_probabilities(&self) -> &[f64] {
245        &self.initial
246    }
247
248    /// Returns the shared finite transition representation.
249    pub fn transitions(&self) -> &FiniteTransitionMatrix<S> {
250        &self.transitions
251    }
252
253    /// Returns the categorical or Gaussian emission representation.
254    pub fn emissions(&self) -> &EmissionModel {
255        &self.emissions
256    }
257
258    pub(crate) fn state_count(&self) -> usize {
259        self.transitions.len()
260    }
261
262    pub(crate) fn emission_log<O: HmmObservation>(
263        &self,
264        state: usize,
265        observation: O,
266    ) -> Result<f64, HmmError> {
267        observation.emission_log_likelihood(&self.emissions, state)
268    }
269}
270
271/// Failure while constructing, fitting, or running hidden-state inference.
272#[derive(Clone, Debug, PartialEq)]
273pub enum HmmError {
274    /// A finite transition matrix was malformed.
275    Transition(TransitionError),
276    /// No observations or fitting sequences were supplied.
277    EmptyInput,
278    /// One fitting sequence contained no observations.
279    EmptySequence {
280        /// Zero-based sequence index.
281        index: usize,
282    },
283    /// The initial or emission state dimension was wrong.
284    EmissionStateCount {
285        /// Required hidden-state count.
286        expected: usize,
287        /// Supplied count.
288        actual: usize,
289    },
290    /// A model-level field was invalid.
291    InvalidModel {
292        /// Invalid field.
293        field: &'static str,
294        /// Concrete requirement.
295        reason: &'static str,
296    },
297    /// A Gaussian parameter was invalid.
298    InvalidGaussian {
299        /// Hidden-state index.
300        state: usize,
301        /// Invalid parameter name.
302        field: &'static str,
303        /// Rejected value.
304        value: f64,
305    },
306    /// Observation and emission kinds did not match.
307    EmissionKind {
308        /// Required observation kind.
309        expected: &'static str,
310        /// Model emission kind.
311        actual: &'static str,
312    },
313    /// A categorical observation was outside the model vocabulary.
314    UnknownSymbol {
315        /// Rejected symbol index.
316        symbol: usize,
317        /// Model symbol count.
318        symbol_count: usize,
319    },
320    /// A continuous observation was not finite.
321    NonFiniteObservation {
322        /// Rejected observation.
323        value: f64,
324    },
325    /// Every state path had zero probability at one observation.
326    ImpossibleSequence {
327        /// Zero-based observation index.
328        position: usize,
329    },
330    /// Fitting data mixed discrete and continuous sequences.
331    MixedSequenceKinds,
332    /// A fitting specification or control was invalid.
333    InvalidFitControl {
334        /// Invalid field.
335        field: &'static str,
336        /// Concrete requirement.
337        reason: &'static str,
338    },
339    /// Checked fitting work arithmetic overflowed.
340    WorkOverflow,
341}
342
343impl fmt::Display for HmmError {
344    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
345        match self {
346            Self::Transition(error) => write!(formatter, "{error}"),
347            Self::EmptyInput => write!(formatter, "HMM inference requires observations"),
348            Self::EmptySequence { index } => {
349                write!(formatter, "HMM fitting sequence {index} is empty")
350            }
351            Self::EmissionStateCount { expected, actual } => write!(
352                formatter,
353                "HMM emissions require {expected} state rows, got {actual}"
354            ),
355            Self::InvalidModel { field, reason } => {
356                write!(formatter, "invalid HMM model {field}: {reason}")
357            }
358            Self::InvalidGaussian {
359                state,
360                field,
361                value,
362            } => write!(
363                formatter,
364                "HMM Gaussian state {state} {field} is invalid: {value}"
365            ),
366            Self::EmissionKind { expected, actual } => write!(
367                formatter,
368                "HMM expects {expected} observations, model emissions are {actual}"
369            ),
370            Self::UnknownSymbol {
371                symbol,
372                symbol_count,
373            } => write!(
374                formatter,
375                "HMM symbol {symbol} is outside vocabulary 0..{symbol_count}"
376            ),
377            Self::NonFiniteObservation { value } => {
378                write!(formatter, "HMM observation is not finite: {value}")
379            }
380            Self::ImpossibleSequence { position } => write!(
381                formatter,
382                "HMM observation {position} has zero probability under every state path"
383            ),
384            Self::MixedSequenceKinds => {
385                write!(formatter, "HMM fitting data mixes observation kinds")
386            }
387            Self::InvalidFitControl { field, reason } => {
388                write!(formatter, "invalid HMM fit control {field}: {reason}")
389            }
390            Self::WorkOverflow => write!(formatter, "HMM fitting work bound overflow"),
391        }
392    }
393}
394
395impl Error for HmmError {}
396
397impl From<TransitionError> for HmmError {
398    fn from(value: TransitionError) -> Self {
399        Self::Transition(value)
400    }
401}
402
403pub(crate) fn log_probability(probability: f64) -> f64 {
404    if probability == 0.0 {
405        f64::NEG_INFINITY
406    } else {
407        probability.ln()
408    }
409}