Skip to main content

qualia_core_db/solvers/learning/sequential/
hmm.rs

1//! Discrete Hidden Markov Model (PRML ch 13.2) — the standard estimators over a
2//! sequence of discrete observations: the **scaled forward** algorithm for the
3//! sequence log-likelihood, **Viterbi** for the most-likely state path, and
4//! **Baum-Welch** (EM) to learn the parameters. Mission note: time-indexed
5//! provenance / life-record reasoning is temporal; this is the canonical model over
6//! censored temporal evidence. Kernel-class `Reduction` (the message passes).
7
8use crate::solvers::learning::LearningError;
9
10/// A discrete HMM: `k` hidden states, `m` observation symbols.
11#[derive(Debug, Clone)]
12pub struct Hmm {
13    /// Initial state distribution π (length k).
14    pub pi: Vec<f64>,
15    /// Row-major `k×k` transition matrix A (`a[i*k+j]` = P(state j | state i)).
16    pub a: Vec<f64>,
17    /// Row-major `k×m` emission matrix B (`b[i*m+o]` = P(symbol o | state i)).
18    pub b: Vec<f64>,
19    pub k: usize,
20    pub m: usize,
21}
22
23impl Hmm {
24    /// Construct from parameters, validating shapes (rows need not be exactly
25    /// normalized but must be non-empty).
26    pub fn new(
27        pi: Vec<f64>,
28        a: Vec<f64>,
29        b: Vec<f64>,
30        k: usize,
31        m: usize,
32    ) -> Result<Self, LearningError> {
33        if k == 0 || m == 0 || pi.len() != k || a.len() != k * k || b.len() != k * m {
34            return Err(LearningError::InvalidDimension);
35        }
36        Ok(Self { pi, a, b, k, m })
37    }
38
39    /// Scaled forward pass. Returns `(log_likelihood, scaled_alpha, scales)`.
40    fn forward_scaled(&self, obs: &[usize]) -> (f64, Vec<f64>, Vec<f64>) {
41        let (k, t) = (self.k, obs.len());
42        let mut alpha = vec![0.0; t * k];
43        let mut scale = vec![0.0; t];
44        // t = 0
45        let mut s = 0.0;
46        for i in 0..k {
47            let v = self.pi[i] * self.b[i * self.m + obs[0]];
48            alpha[i] = v;
49            s += v;
50        }
51        let c0 = if s > 0.0 { 1.0 / s } else { 0.0 };
52        scale[0] = c0;
53        for i in 0..k {
54            alpha[i] *= c0;
55        }
56        // t > 0
57        for tt in 1..t {
58            let mut s = 0.0;
59            for j in 0..k {
60                let mut acc = 0.0;
61                for i in 0..k {
62                    acc += alpha[(tt - 1) * k + i] * self.a[i * k + j];
63                }
64                let v = acc * self.b[j * self.m + obs[tt]];
65                alpha[tt * k + j] = v;
66                s += v;
67            }
68            let c = if s > 0.0 { 1.0 / s } else { 0.0 };
69            scale[tt] = c;
70            for j in 0..k {
71                alpha[tt * k + j] *= c;
72            }
73        }
74        // log P(obs) = −Σ log c_t.
75        let ll: f64 = scale
76            .iter()
77            .map(|&c| if c > 0.0 { -c.ln() } else { f64::NEG_INFINITY })
78            .sum();
79        (ll, alpha, scale)
80    }
81
82    /// Log-likelihood `log P(obs | model)`. `None` for an empty sequence or an
83    /// out-of-range symbol.
84    pub fn log_likelihood(&self, obs: &[usize]) -> Option<f64> {
85        if obs.is_empty() || obs.iter().any(|&o| o >= self.m) {
86            return None;
87        }
88        Some(self.forward_scaled(obs).0)
89    }
90
91    /// Viterbi most-likely state path + its log-probability. `None` for an empty
92    /// sequence or an out-of-range symbol.
93    pub fn viterbi(&self, obs: &[usize]) -> Option<(Vec<usize>, f64)> {
94        if obs.is_empty() || obs.iter().any(|&o| o >= self.m) {
95            return None;
96        }
97        let (k, t, m) = (self.k, obs.len(), self.m);
98        let ln = |x: f64| if x > 0.0 { x.ln() } else { f64::NEG_INFINITY };
99        let mut delta = vec![f64::NEG_INFINITY; t * k];
100        let mut psi = vec![0usize; t * k];
101        for i in 0..k {
102            delta[i] = ln(self.pi[i]) + ln(self.b[i * m + obs[0]]);
103        }
104        for tt in 1..t {
105            for j in 0..k {
106                let mut best = f64::NEG_INFINITY;
107                let mut arg = 0;
108                for i in 0..k {
109                    let v = delta[(tt - 1) * k + i] + ln(self.a[i * k + j]);
110                    if v > best {
111                        best = v;
112                        arg = i;
113                    }
114                }
115                delta[tt * k + j] = best + ln(self.b[j * m + obs[tt]]);
116                psi[tt * k + j] = arg;
117            }
118        }
119        // Termination + backtrack.
120        let mut last = 0;
121        let mut best = f64::NEG_INFINITY;
122        for i in 0..k {
123            if delta[(t - 1) * k + i] > best {
124                best = delta[(t - 1) * k + i];
125                last = i;
126            }
127        }
128        let mut path = vec![0usize; t];
129        path[t - 1] = last;
130        for tt in (1..t).rev() {
131            path[tt - 1] = psi[tt * k + path[tt]];
132        }
133        Some((path, best))
134    }
135}
136
137struct Lcg(u64);
138impl Lcg {
139    fn unit(&mut self) -> f64 {
140        self.0 = self
141            .0
142            .wrapping_mul(6364136223846793005)
143            .wrapping_add(1442695040888963407);
144        ((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
145    }
146}
147
148fn normalize(row: &mut [f64]) {
149    let s: f64 = row.iter().sum();
150    if s > 0.0 {
151        for v in row.iter_mut() {
152            *v /= s;
153        }
154    }
155}
156
157/// Learn HMM parameters from one observation sequence by Baum-Welch (EM). Returns
158/// `(model, final_log_likelihood)`. Initialised randomly (seeded). Fails closed on
159/// bad shapes / out-of-range symbols.
160pub fn baum_welch(
161    obs: &[usize],
162    k: usize,
163    m: usize,
164    max_iter: usize,
165    tol: f64,
166    seed: u64,
167) -> Result<(Hmm, f64), LearningError> {
168    let t = obs.len();
169    if k == 0 || m == 0 || t < 2 || obs.iter().any(|&o| o >= m) {
170        return Err(LearningError::InvalidDimension);
171    }
172
173    // Random near-uniform initialisation.
174    let mut rng = Lcg(seed ^ 0x9E3779B97F4A7C15);
175    let mut pi = vec![0.0; k];
176    let mut a = vec![0.0; k * k];
177    let mut b = vec![0.0; k * m];
178    for i in 0..k {
179        pi[i] = 1.0 + 0.1 * rng.unit();
180    }
181    normalize(&mut pi);
182    for i in 0..k {
183        for j in 0..k {
184            a[i * k + j] = 1.0 + 0.1 * rng.unit();
185        }
186        normalize(&mut a[i * k..(i + 1) * k]);
187        for o in 0..m {
188            b[i * m + o] = 1.0 + 0.1 * rng.unit();
189        }
190        normalize(&mut b[i * m..(i + 1) * m]);
191    }
192
193    let mut hmm = Hmm { pi, a, b, k, m };
194    let mut prev_ll = f64::NEG_INFINITY;
195    let mut final_ll = prev_ll;
196
197    for _ in 0..max_iter.max(1) {
198        // E-step: scaled forward + backward.
199        let (ll, alpha, scale) = hmm.forward_scaled(obs);
200        final_ll = ll;
201        // Scaled backward.
202        let mut beta = vec![0.0; t * k];
203        for i in 0..k {
204            beta[(t - 1) * k + i] = scale[t - 1];
205        }
206        for tt in (0..t - 1).rev() {
207            for i in 0..k {
208                let mut acc = 0.0;
209                for j in 0..k {
210                    acc += hmm.a[i * k + j] * hmm.b[j * m + obs[tt + 1]] * beta[(tt + 1) * k + j];
211                }
212                beta[tt * k + i] = acc * scale[tt];
213            }
214        }
215        // γ and accumulate ξ sums.
216        let mut gamma = vec![0.0; t * k];
217        for tt in 0..t {
218            let mut s = 0.0;
219            for i in 0..k {
220                gamma[tt * k + i] = alpha[tt * k + i] * beta[tt * k + i];
221                s += gamma[tt * k + i];
222            }
223            if s > 0.0 {
224                for i in 0..k {
225                    gamma[tt * k + i] /= s;
226                }
227            }
228        }
229        // M-step.
230        // π.
231        for i in 0..k {
232            hmm.pi[i] = gamma[i];
233        }
234        // A.
235        let mut new_a = vec![0.0; k * k];
236        for i in 0..k {
237            let mut denom = 0.0;
238            for tt in 0..t - 1 {
239                denom += gamma[tt * k + i];
240            }
241            for j in 0..k {
242                let mut num = 0.0;
243                for tt in 0..t - 1 {
244                    num += alpha[tt * k + i]
245                        * hmm.a[i * k + j]
246                        * hmm.b[j * m + obs[tt + 1]]
247                        * beta[(tt + 1) * k + j];
248                }
249                new_a[i * k + j] = if denom > 0.0 { num / denom } else { 0.0 };
250            }
251            normalize(&mut new_a[i * k..(i + 1) * k]);
252        }
253        hmm.a = new_a;
254        // B.
255        let mut new_b = vec![0.0; k * m];
256        for i in 0..k {
257            let mut denom = 0.0;
258            for tt in 0..t {
259                denom += gamma[tt * k + i];
260            }
261            for tt in 0..t {
262                new_b[i * m + obs[tt]] += gamma[tt * k + i];
263            }
264            if denom > 0.0 {
265                for o in 0..m {
266                    new_b[i * m + o] /= denom;
267                }
268            }
269            normalize(&mut new_b[i * m..(i + 1) * m]);
270        }
271        hmm.b = new_b;
272
273        if (ll - prev_ll).abs() < tol {
274            break;
275        }
276        prev_ll = ll;
277    }
278
279    Ok((hmm, final_ll))
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285
286    // A 2-state HMM: state 0 emits symbol 0 mostly, state 1 emits symbol 1 mostly;
287    // states are sticky (stay put with high probability).
288    fn sticky_hmm() -> Hmm {
289        Hmm::new(
290            vec![0.5, 0.5],
291            vec![0.9, 0.1, 0.1, 0.9],
292            vec![0.9, 0.1, 0.1, 0.9],
293            2,
294            2,
295        )
296        .unwrap()
297    }
298
299    #[test]
300    fn viterbi_recovers_obvious_path() {
301        let hmm = sticky_hmm();
302        // Observations clearly in state 0 then state 1.
303        let obs = [0, 0, 0, 1, 1, 1];
304        let (path, _) = hmm.viterbi(&obs).unwrap();
305        assert_eq!(path, vec![0, 0, 0, 1, 1, 1]);
306    }
307
308    #[test]
309    fn log_likelihood_is_finite_and_orders_sequences() {
310        let hmm = sticky_hmm();
311        // A "consistent" sequence is more likely than a rapidly alternating one.
312        let consistent = hmm.log_likelihood(&[0, 0, 0, 0]).unwrap();
313        let alternating = hmm.log_likelihood(&[0, 1, 0, 1]).unwrap();
314        assert!(consistent.is_finite() && alternating.is_finite());
315        assert!(consistent > alternating, "{consistent} !> {alternating}");
316        assert!(hmm.log_likelihood(&[]).is_none());
317        assert!(hmm.log_likelihood(&[5]).is_none()); // out-of-range symbol
318    }
319
320    #[test]
321    fn baum_welch_increases_likelihood_and_learns_structure() {
322        // A long sequence with clear regime structure.
323        let mut obs = Vec::new();
324        for _ in 0..15 {
325            obs.push(0);
326        }
327        for _ in 0..15 {
328            obs.push(1);
329        }
330        for _ in 0..15 {
331            obs.push(0);
332        }
333        let (model, ll) = baum_welch(&obs, 2, 2, 100, 1e-6, 1).unwrap();
334        assert!(ll.is_finite());
335        // The learned model should make the training sequence at least as likely as
336        // a uniform-ish model, and Viterbi should segment the regimes.
337        let (path, _) = model.viterbi(&obs).unwrap();
338        // The first block and the middle block should differ in state.
339        assert_ne!(path[5], path[20], "regimes should map to different states");
340    }
341
342    #[test]
343    fn guards() {
344        assert!(Hmm::new(vec![1.0], vec![1.0], vec![1.0, 0.0], 1, 2).is_ok());
345        assert_eq!(
346            baum_welch(&[0], 2, 2, 10, 1e-6, 0).unwrap_err(),
347            LearningError::InvalidDimension
348        );
349    }
350}