qualia_core_db/solvers/learning/sequential/
hmm.rs1use crate::solvers::learning::LearningError;
9
10#[derive(Debug, Clone)]
12pub struct Hmm {
13 pub pi: Vec<f64>,
15 pub a: Vec<f64>,
17 pub b: Vec<f64>,
19 pub k: usize,
20 pub m: usize,
21}
22
23impl Hmm {
24 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 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 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 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 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 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 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 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
157pub 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 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 let (ll, alpha, scale) = hmm.forward_scaled(obs);
200 final_ll = ll;
201 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 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 for i in 0..k {
232 hmm.pi[i] = gamma[i];
233 }
234 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 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 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 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 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()); }
319
320 #[test]
321 fn baum_welch_increases_likelihood_and_learns_structure() {
322 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 let (path, _) = model.viterbi(&obs).unwrap();
338 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}