1use crate::regimes::MarketRegime;
12use crate::traits::Next;
13use serde::{Deserialize, Serialize};
14
15const VAR_FLOOR: f64 = 1e-9;
16const LOG_FLOOR: f64 = 1e-300;
17
18#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct GMM {
21 k: usize,
22 dims: usize,
23 means: Vec<Vec<f64>>,
25 vars: Vec<Vec<f64>>,
27 weights: Vec<f64>,
29}
30
31#[derive(Debug, Clone)]
33pub struct GmmFitConfig {
34 pub max_iter: usize,
35 pub tol: f64,
36 pub seed: u64,
37}
38
39impl Default for GmmFitConfig {
40 fn default() -> Self {
41 Self {
42 max_iter: 100,
43 tol: 1e-6,
44 seed: 42,
45 }
46 }
47}
48
49#[derive(Debug, Clone)]
51pub struct GmmFitResult {
52 pub log_likelihood: f64,
53 pub iterations: usize,
54 pub converged: bool,
55}
56
57#[derive(Debug, thiserror::Error, PartialEq)]
58pub enum GmmError {
59 #[error("invalid GMM parameters: {0}")]
60 InvalidParams(String),
61 #[error("need at least {min} observations, got {got}")]
62 InsufficientData { min: usize, got: usize },
63 #[error("EM did not converge within {max_iter} iterations")]
64 EmNotConverged { max_iter: usize },
65}
66
67impl GMM {
68 pub fn new(means: Vec<Vec<f64>>, vars: Vec<Vec<f64>>, weights: Vec<f64>) -> Self {
70 let k = means.len();
71 let dims = means[0].len();
72 Self {
73 k,
74 dims,
75 means,
76 vars,
77 weights,
78 }
79 }
80
81 pub fn with_components(k: usize, dims: usize) -> Self {
83 let means = vec![vec![0.0; dims]; k];
84 let vars = vec![vec![1.0; dims]; k];
85 let weights = vec![1.0 / k as f64; k];
86 Self {
87 k,
88 dims,
89 means,
90 vars,
91 weights,
92 }
93 }
94
95 pub fn components(&self) -> usize {
96 self.k
97 }
98
99 pub fn dims(&self) -> usize {
100 self.dims
101 }
102
103 pub fn means(&self) -> &[Vec<f64>] {
104 &self.means
105 }
106
107 pub fn weights(&self) -> &[f64] {
108 &self.weights
109 }
110
111 fn log_pdf(&self, x: &[f64], k_idx: usize) -> f64 {
113 let mut log_prob = 0.0;
114 for d in 0..self.dims {
115 let mu = self.means[k_idx][d];
116 let var = self.vars[k_idx][d].max(VAR_FLOOR);
117 let diff = x[d] - mu;
118 log_prob += -0.5 * ((2.0 * std::f64::consts::PI * var).ln() + diff * diff / var);
119 }
120 log_prob
121 }
122
123 fn pdf(&self, x: &[f64], k_idx: usize) -> f64 {
125 self.log_pdf(x, k_idx).exp()
126 }
127
128 fn validate_data(&self, data: &[Vec<f64>]) -> Result<(), GmmError> {
129 if data.len() < self.k {
130 return Err(GmmError::InsufficientData {
131 min: self.k,
132 got: data.len(),
133 });
134 }
135 for row in data {
136 if row.len() != self.dims {
137 return Err(GmmError::InvalidParams(format!(
138 "expected {dims} dims, got {got}",
139 dims = self.dims,
140 got = row.len()
141 )));
142 }
143 }
144 Ok(())
145 }
146
147 fn init_from_quantiles(&mut self, data: &[Vec<f64>]) {
148 let n = data.len();
149 let mut order: Vec<usize> = (0..n).collect();
150 order.sort_by(|&a, &b| {
151 data[a][0]
152 .partial_cmp(&data[b][0])
153 .unwrap_or(std::cmp::Ordering::Equal)
154 });
155
156 for (k, chunk) in order.chunks((n / self.k).max(1)).enumerate().take(self.k) {
157 if chunk.is_empty() {
158 continue;
159 }
160 for d in 0..self.dims {
161 let sum: f64 = chunk.iter().map(|&i| data[i][d]).sum();
162 self.means[k][d] = sum / chunk.len() as f64;
163 let var: f64 = chunk
164 .iter()
165 .map(|&i| {
166 let diff = data[i][d] - self.means[k][d];
167 diff * diff
168 })
169 .sum::<f64>()
170 / chunk.len() as f64;
171 self.vars[k][d] = var.max(VAR_FLOOR);
172 }
173 self.weights[k] = chunk.len() as f64 / n as f64;
174 }
175
176 let w_sum: f64 = self.weights.iter().sum();
177 if w_sum > 0.0 {
178 for w in &mut self.weights {
179 *w /= w_sum;
180 }
181 }
182 }
183
184 fn responsibilities(&self, data: &[Vec<f64>]) -> Vec<Vec<f64>> {
185 let n = data.len();
186 let mut resp = vec![vec![0.0; self.k]; n];
187 for (i, x) in data.iter().enumerate() {
188 let mut log_probs = vec![0.0; self.k];
189 let mut max_log = f64::NEG_INFINITY;
190 for k in 0..self.k {
191 let lp = self.weights[k].max(LOG_FLOOR).ln() + self.log_pdf(x, k);
192 log_probs[k] = lp;
193 if lp > max_log {
194 max_log = lp;
195 }
196 }
197 let mut sum = 0.0;
198 for k in 0..self.k {
199 let r = (log_probs[k] - max_log).exp();
200 resp[i][k] = r;
201 sum += r;
202 }
203 if sum > 0.0 {
204 for k in 0..self.k {
205 resp[i][k] /= sum;
206 }
207 }
208 }
209 resp
210 }
211
212 fn log_likelihood(&self, data: &[Vec<f64>]) -> f64 {
213 let mut total = 0.0;
214 for x in data {
215 let mut log_probs = vec![0.0; self.k];
216 let mut max_log = f64::NEG_INFINITY;
217 for k in 0..self.k {
218 let lp = self.weights[k].max(LOG_FLOOR).ln() + self.log_pdf(x, k);
219 log_probs[k] = lp;
220 if lp > max_log {
221 max_log = lp;
222 }
223 }
224 let ll = max_log
225 + log_probs
226 .iter()
227 .map(|&lp| (lp - max_log).exp())
228 .sum::<f64>()
229 .ln();
230 total += ll;
231 }
232 total
233 }
234
235 fn m_step(&mut self, data: &[Vec<f64>], resp: &[Vec<f64>]) {
236 let n = data.len();
237 for k in 0..self.k {
238 let nk: f64 = resp.iter().map(|r| r[k]).sum();
239 if nk < LOG_FLOOR {
240 continue;
241 }
242 self.weights[k] = nk / n as f64;
243 for d in 0..self.dims {
244 let mean: f64 = resp
245 .iter()
246 .zip(data.iter())
247 .map(|(r, x)| r[k] * x[d])
248 .sum::<f64>()
249 / nk;
250 self.means[k][d] = mean;
251 let var: f64 = resp
252 .iter()
253 .zip(data.iter())
254 .map(|(r, x)| {
255 let diff = x[d] - mean;
256 r[k] * diff * diff
257 })
258 .sum::<f64>()
259 / nk;
260 self.vars[k][d] = var.max(VAR_FLOOR);
261 }
262 }
263 }
264
265 pub fn fit(
267 &mut self,
268 data: &[Vec<f64>],
269 config: &GmmFitConfig,
270 ) -> Result<GmmFitResult, GmmError> {
271 self.validate_data(data)?;
272 self.init_from_quantiles(data);
273
274 let mut prev_ll = f64::NEG_INFINITY;
275 let mut iterations = 0usize;
276 let mut converged = false;
277
278 for iter in 0..config.max_iter {
279 iterations = iter + 1;
280 let resp = self.responsibilities(data);
281 self.m_step(data, &resp);
282 let ll = self.log_likelihood(data);
283 if (ll - prev_ll).abs() < config.tol {
284 converged = true;
285 prev_ll = ll;
286 break;
287 }
288 if ll < prev_ll - config.tol {
289 }
291 prev_ll = ll;
292 }
293
294 Ok(GmmFitResult {
295 log_likelihood: prev_ll,
296 iterations,
297 converged,
298 })
299 }
300}
301
302impl Next<&[f64]> for GMM {
303 type Output = MarketRegime;
304
305 fn next(&mut self, x: &[f64]) -> Self::Output {
306 let mut max_prob = -1.0;
307 let mut best_k = 0;
308
309 for k in 0..self.k {
310 let p = self.weights[k] * self.pdf(x, k);
311 if p > max_prob {
312 max_prob = p;
313 best_k = k;
314 }
315 }
316
317 match best_k {
318 0 => MarketRegime::Steady,
319 k if k == self.k - 1 => MarketRegime::Crisis,
320 _ => MarketRegime::Cluster(best_k as u8),
321 }
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328 use approx::assert_relative_eq;
329
330 fn sample_three_gaussians(seed: u64) -> (Vec<Vec<f64>>, Vec<f64>) {
331 let mut data = Vec::new();
332 let true_means = [-5.0, 0.0, 5.0];
333 let mut state = seed;
334 for (c, &mu) in true_means.iter().enumerate() {
335 for _ in 0..200 {
336 state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
337 let u = (state >> 11) as f64 / (1u64 << 53) as f64;
338 let v = (state >> 17) as f64 / (1u64 << 47) as f64;
339 let z = (-2.0 * u.ln()).sqrt() * (2.0 * std::f64::consts::PI * v).cos();
340 data.push(vec![mu + z * 0.5]);
341 let _ = c;
342 }
343 }
344 (data, true_means.to_vec())
345 }
346
347 #[test]
348 fn fit_recovers_three_gaussian_means() {
349 let (data, true_means) = sample_three_gaussians(99);
350 let mut gmm = GMM::with_components(3, 1);
351 let result = gmm
352 .fit(&data, &GmmFitConfig::default())
353 .expect("fit should succeed");
354 assert!(result.converged);
355 let mut recovered: Vec<f64> = gmm.means().iter().map(|m| m[0]).collect();
356 recovered.sort_by(|a, b| a.partial_cmp(b).unwrap());
357 let mut expected = true_means;
358 expected.sort_by(|a, b| a.partial_cmp(b).unwrap());
359 for (r, e) in recovered.iter().zip(expected.iter()) {
360 assert_relative_eq!(r, e, epsilon = 0.75);
361 }
362 for w in gmm.weights() {
363 assert_relative_eq!(*w, 1.0 / 3.0, epsilon = 0.15);
364 }
365 }
366
367 #[test]
368 fn fit_insufficient_data_errors() {
369 let mut gmm = GMM::with_components(3, 1);
370 let err = gmm.fit(&[vec![1.0], vec![2.0]], &GmmFitConfig::default());
371 assert!(matches!(err, Err(GmmError::InsufficientData { .. })));
372 }
373
374 #[test]
375 fn log_likelihood_non_decreasing_on_easy_data() {
376 let (data, _) = sample_three_gaussians(7);
377 let mut gmm = GMM::with_components(3, 1);
378 gmm.validate_data(&data).unwrap();
379 gmm.init_from_quantiles(&data);
380 let mut prev = f64::NEG_INFINITY;
381 for _ in 0..10 {
382 let resp = gmm.responsibilities(&data);
383 gmm.m_step(&data, &resp);
384 let ll = gmm.log_likelihood(&data);
385 assert!(ll >= prev - 1e-9, "LL decreased: {prev} -> {ll}");
386 prev = ll;
387 }
388 }
389
390 #[test]
391 fn max_iter_one_reports_not_converged() {
392 let (data, _) = sample_three_gaussians(3);
393 let mut gmm = GMM::with_components(3, 1);
394 let cfg = GmmFitConfig {
395 max_iter: 1,
396 tol: 1e-12,
397 seed: 1,
398 };
399 let result = gmm.fit(&data, &cfg).unwrap();
400 assert!(!result.converged);
401 assert_eq!(result.iterations, 1);
402 }
403}