Skip to main content

sim_lib_numbers_stats/
hmm_inference.rs

1//! Underflow-safe normalized forward/backward and path inference.
2
3use super::hmm_model::{HiddenMarkovModel, HmmError, HmmObservation, log_probability};
4
5/// Numerical and likelihood evidence from normalized sequence inference.
6#[derive(Clone, Copy, Debug, PartialEq)]
7pub struct InferenceEvidence {
8    /// Natural logarithm of the full observation-sequence likelihood.
9    pub log_likelihood: f64,
10    /// Count of finite renormalizations needed beyond log-domain scaling.
11    pub numerical_repairs: u64,
12    /// Number of normalized time steps.
13    pub normalized_steps: usize,
14}
15
16/// Normalized forward, backward, and smoothed posterior state probabilities.
17#[derive(Clone, Debug, PartialEq)]
18pub struct ForwardBackward {
19    /// Time-major normalized filtering probabilities.
20    pub forward: Vec<Vec<f64>>,
21    /// Time-major normalized backward likelihood weights.
22    pub backward: Vec<Vec<f64>>,
23    /// Time-major normalized smoothed state probabilities.
24    pub posterior: Vec<Vec<f64>>,
25    /// Likelihood and normalization diagnostics.
26    pub evidence: InferenceEvidence,
27}
28
29/// Maximum-probability hidden-state path and its joint log probability.
30#[derive(Clone, Debug, PartialEq)]
31pub struct ViterbiPath<S> {
32    /// Hidden states in observation order.
33    pub states: Vec<S>,
34    /// State indices in the model's stable order.
35    pub state_indices: Vec<usize>,
36    /// Natural logarithm of the path and observations' joint probability.
37    pub log_probability: f64,
38    /// Number of numerical repairs; log-domain Viterbi normally reports zero.
39    pub numerical_repairs: u64,
40}
41
42/// Per-position maximum-posterior state path.
43#[derive(Clone, Debug, PartialEq)]
44pub struct PosteriorPath<S> {
45    /// Hidden states in observation order.
46    pub states: Vec<S>,
47    /// State indices in the model's stable order.
48    pub state_indices: Vec<usize>,
49    /// Posterior probability of each selected state.
50    pub confidence: Vec<f64>,
51    /// Likelihood and normalization diagnostics.
52    pub evidence: InferenceEvidence,
53}
54
55/// Runs normalized forward/backward inference in the log domain.
56pub fn forward_backward<O: HmmObservation, S: Eq + Clone>(
57    model: &HiddenMarkovModel<S>,
58    observations: &[O],
59) -> Result<ForwardBackward, HmmError> {
60    if observations.is_empty() {
61        return Err(HmmError::EmptyInput);
62    }
63    let states = model.state_count();
64    let mut repairs = 0_u64;
65    let mut forward = Vec::with_capacity(observations.len());
66    let mut first = (0..states)
67        .map(|state| {
68            Ok(log_probability(model.initial_probabilities()[state])
69                + model.emission_log(state, observations[0])?)
70        })
71        .collect::<Result<Vec<_>, HmmError>>()?;
72    let mut log_likelihood = normalize_logs(&mut first, 0, &mut repairs)?;
73    forward.push(first);
74
75    for (position, observation) in observations.iter().copied().enumerate().skip(1) {
76        let previous = forward.last().expect("forward row exists");
77        let mut row = Vec::with_capacity(states);
78        for to in 0..states {
79            let terms = (0..states).map(|from| {
80                log_probability(previous[from])
81                    + log_probability(
82                        model
83                            .transitions()
84                            .probability_by_index(from, to)
85                            .unwrap_or(0.0),
86                    )
87            });
88            row.push(log_sum_exp(terms) + model.emission_log(to, observation)?);
89        }
90        log_likelihood += normalize_logs(&mut row, position, &mut repairs)?;
91        forward.push(row);
92    }
93
94    let mut backward = vec![vec![0.0; states]; observations.len()];
95    backward[observations.len() - 1].fill(1.0 / states as f64);
96    for position in (0..observations.len() - 1).rev() {
97        let mut row = Vec::with_capacity(states);
98        for from in 0..states {
99            let mut terms = Vec::with_capacity(states);
100            for (to, &backward_probability) in backward[position + 1].iter().enumerate() {
101                terms.push(
102                    log_probability(
103                        model
104                            .transitions()
105                            .probability_by_index(from, to)
106                            .unwrap_or(0.0),
107                    ) + model.emission_log(to, observations[position + 1])?
108                        + log_probability(backward_probability),
109                );
110            }
111            row.push(log_sum_exp(terms));
112        }
113        normalize_logs(&mut row, position, &mut repairs)?;
114        backward[position] = row;
115    }
116
117    let posterior = forward
118        .iter()
119        .zip(&backward)
120        .enumerate()
121        .map(|(position, (alpha, beta))| {
122            let mut row = alpha
123                .iter()
124                .zip(beta)
125                .map(|(alpha, beta)| alpha * beta)
126                .collect::<Vec<_>>();
127            normalize_weights(&mut row, position)?;
128            Ok(row)
129        })
130        .collect::<Result<Vec<_>, HmmError>>()?;
131    Ok(ForwardBackward {
132        forward,
133        backward,
134        posterior,
135        evidence: InferenceEvidence {
136            log_likelihood,
137            numerical_repairs: repairs,
138            normalized_steps: observations.len(),
139        },
140    })
141}
142
143/// Finds the maximum joint-probability hidden-state path in the log domain.
144pub fn viterbi<O: HmmObservation, S: Eq + Clone>(
145    model: &HiddenMarkovModel<S>,
146    observations: &[O],
147) -> Result<ViterbiPath<S>, HmmError> {
148    if observations.is_empty() {
149        return Err(HmmError::EmptyInput);
150    }
151    let states = model.state_count();
152    let mut scores = (0..states)
153        .map(|state| {
154            Ok(log_probability(model.initial_probabilities()[state])
155                + model.emission_log(state, observations[0])?)
156        })
157        .collect::<Result<Vec<_>, HmmError>>()?;
158    require_possible(&scores, 0)?;
159    let mut backpointers = Vec::with_capacity(observations.len().saturating_sub(1));
160    for (position, observation) in observations.iter().copied().enumerate().skip(1) {
161        let mut next = vec![f64::NEG_INFINITY; states];
162        let mut pointers = vec![0; states];
163        for to in 0..states {
164            for (from, &score) in scores.iter().enumerate() {
165                let candidate = score
166                    + log_probability(
167                        model
168                            .transitions()
169                            .probability_by_index(from, to)
170                            .unwrap_or(0.0),
171                    );
172                if candidate > next[to] {
173                    next[to] = candidate;
174                    pointers[to] = from;
175                }
176            }
177            next[to] += model.emission_log(to, observation)?;
178        }
179        require_possible(&next, position)?;
180        scores = next;
181        backpointers.push(pointers);
182    }
183    let (mut state, &log_probability) = scores
184        .iter()
185        .enumerate()
186        .max_by(|(left_index, left), (right_index, right)| {
187            left.total_cmp(right)
188                .then_with(|| right_index.cmp(left_index))
189        })
190        .expect("non-empty hidden state set");
191    let mut state_indices = vec![state];
192    for pointers in backpointers.iter().rev() {
193        state = pointers[state];
194        state_indices.push(state);
195    }
196    state_indices.reverse();
197    let states = state_indices
198        .iter()
199        .map(|&index| model.states()[index].clone())
200        .collect();
201    Ok(ViterbiPath {
202        states,
203        state_indices,
204        log_probability,
205        numerical_repairs: 0,
206    })
207}
208
209/// Decodes the independently most probable hidden state at each position.
210pub fn posterior_decode<O: HmmObservation, S: Eq + Clone>(
211    model: &HiddenMarkovModel<S>,
212    observations: &[O],
213) -> Result<PosteriorPath<S>, HmmError> {
214    let inference = forward_backward(model, observations)?;
215    let mut state_indices = Vec::with_capacity(observations.len());
216    let mut confidence = Vec::with_capacity(observations.len());
217    for row in &inference.posterior {
218        let (state, &probability) = row
219            .iter()
220            .enumerate()
221            .max_by(|(left_index, left), (right_index, right)| {
222                left.total_cmp(right)
223                    .then_with(|| right_index.cmp(left_index))
224            })
225            .expect("non-empty hidden state set");
226        state_indices.push(state);
227        confidence.push(probability);
228    }
229    let states = state_indices
230        .iter()
231        .map(|&index| model.states()[index].clone())
232        .collect();
233    Ok(PosteriorPath {
234        states,
235        state_indices,
236        confidence,
237        evidence: inference.evidence,
238    })
239}
240
241pub(crate) fn log_sum_exp(values: impl IntoIterator<Item = f64>) -> f64 {
242    let values = values.into_iter().collect::<Vec<_>>();
243    let maximum = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
244    if maximum == f64::NEG_INFINITY {
245        return maximum;
246    }
247    maximum
248        + values
249            .iter()
250            .map(|value| (value - maximum).exp())
251            .sum::<f64>()
252            .ln()
253}
254
255fn normalize_logs(values: &mut [f64], position: usize, repairs: &mut u64) -> Result<f64, HmmError> {
256    let log_total = log_sum_exp(values.iter().copied());
257    if !log_total.is_finite() {
258        return Err(HmmError::ImpossibleSequence { position });
259    }
260    for value in values.iter_mut() {
261        *value = (*value - log_total).exp();
262    }
263    normalize_probabilities(values, position, repairs)?;
264    Ok(log_total)
265}
266
267fn normalize_probabilities(
268    values: &mut [f64],
269    position: usize,
270    repairs: &mut u64,
271) -> Result<(), HmmError> {
272    let sum = values.iter().sum::<f64>();
273    if !sum.is_finite() || sum <= 0.0 {
274        return Err(HmmError::ImpossibleSequence { position });
275    }
276    if (sum - 1.0).abs() > 1.0e-12 {
277        *repairs = repairs.saturating_add(1);
278    }
279    for value in values {
280        *value /= sum;
281    }
282    Ok(())
283}
284
285fn normalize_weights(values: &mut [f64], position: usize) -> Result<(), HmmError> {
286    let sum = values.iter().sum::<f64>();
287    if !sum.is_finite() || sum <= 0.0 {
288        return Err(HmmError::ImpossibleSequence { position });
289    }
290    for value in values {
291        *value /= sum;
292    }
293    Ok(())
294}
295
296fn require_possible(scores: &[f64], position: usize) -> Result<(), HmmError> {
297    if scores.iter().any(|score| score.is_finite()) {
298        Ok(())
299    } else {
300        Err(HmmError::ImpossibleSequence { position })
301    }
302}