1use super::hmm_model::{HiddenMarkovModel, HmmError, HmmObservation, log_probability};
4
5#[derive(Clone, Copy, Debug, PartialEq)]
7pub struct InferenceEvidence {
8 pub log_likelihood: f64,
10 pub numerical_repairs: u64,
12 pub normalized_steps: usize,
14}
15
16#[derive(Clone, Debug, PartialEq)]
18pub struct ForwardBackward {
19 pub forward: Vec<Vec<f64>>,
21 pub backward: Vec<Vec<f64>>,
23 pub posterior: Vec<Vec<f64>>,
25 pub evidence: InferenceEvidence,
27}
28
29#[derive(Clone, Debug, PartialEq)]
31pub struct ViterbiPath<S> {
32 pub states: Vec<S>,
34 pub state_indices: Vec<usize>,
36 pub log_probability: f64,
38 pub numerical_repairs: u64,
40}
41
42#[derive(Clone, Debug, PartialEq)]
44pub struct PosteriorPath<S> {
45 pub states: Vec<S>,
47 pub state_indices: Vec<usize>,
49 pub confidence: Vec<f64>,
51 pub evidence: InferenceEvidence,
53}
54
55pub 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
143pub 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
209pub 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}