1use super::transition::{FiniteTransitionMatrix, TransitionError, validate_distribution};
4use std::{error::Error, f64::consts::TAU, fmt};
5
6#[derive(Clone, Debug, PartialEq)]
8pub enum EmissionModel {
9 Discrete {
11 probabilities: Vec<Vec<f64>>,
13 },
14 Gaussian {
16 means: Vec<f64>,
18 variances: Vec<f64>,
20 variance_floor: f64,
22 },
23}
24
25impl EmissionModel {
26 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 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 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
147pub trait HmmObservation: Copy {
149 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#[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 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 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 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 pub fn states(&self) -> &[S] {
240 self.transitions.states()
241 }
242
243 pub fn initial_probabilities(&self) -> &[f64] {
245 &self.initial
246 }
247
248 pub fn transitions(&self) -> &FiniteTransitionMatrix<S> {
250 &self.transitions
251 }
252
253 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#[derive(Clone, Debug, PartialEq)]
273pub enum HmmError {
274 Transition(TransitionError),
276 EmptyInput,
278 EmptySequence {
280 index: usize,
282 },
283 EmissionStateCount {
285 expected: usize,
287 actual: usize,
289 },
290 InvalidModel {
292 field: &'static str,
294 reason: &'static str,
296 },
297 InvalidGaussian {
299 state: usize,
301 field: &'static str,
303 value: f64,
305 },
306 EmissionKind {
308 expected: &'static str,
310 actual: &'static str,
312 },
313 UnknownSymbol {
315 symbol: usize,
317 symbol_count: usize,
319 },
320 NonFiniteObservation {
322 value: f64,
324 },
325 ImpossibleSequence {
327 position: usize,
329 },
330 MixedSequenceKinds,
332 InvalidFitControl {
334 field: &'static str,
336 reason: &'static str,
338 },
339 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}