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 + log_probs.iter().map(|&lp| (lp - max_log).exp()).sum::<f64>().ln();
225 total += ll;
226 }
227 total
228 }
229
230 fn m_step(&mut self, data: &[Vec<f64>], resp: &[Vec<f64>]) {
231 let n = data.len();
232 for k in 0..self.k {
233 let nk: f64 = resp.iter().map(|r| r[k]).sum();
234 if nk < LOG_FLOOR {
235 continue;
236 }
237 self.weights[k] = nk / n as f64;
238 for d in 0..self.dims {
239 let mean: f64 = resp
240 .iter()
241 .zip(data.iter())
242 .map(|(r, x)| r[k] * x[d])
243 .sum::<f64>()
244 / nk;
245 self.means[k][d] = mean;
246 let var: f64 = resp
247 .iter()
248 .zip(data.iter())
249 .map(|(r, x)| {
250 let diff = x[d] - mean;
251 r[k] * diff * diff
252 })
253 .sum::<f64>()
254 / nk;
255 self.vars[k][d] = var.max(VAR_FLOOR);
256 }
257 }
258 }
259
260 pub fn fit(
262 &mut self,
263 data: &[Vec<f64>],
264 config: &GmmFitConfig,
265 ) -> Result<GmmFitResult, GmmError> {
266 self.validate_data(data)?;
267 self.init_from_quantiles(data);
268
269 let mut prev_ll = f64::NEG_INFINITY;
270 let mut iterations = 0usize;
271 let mut converged = false;
272
273 for iter in 0..config.max_iter {
274 iterations = iter + 1;
275 let resp = self.responsibilities(data);
276 self.m_step(data, &resp);
277 let ll = self.log_likelihood(data);
278 if (ll - prev_ll).abs() < config.tol {
279 converged = true;
280 prev_ll = ll;
281 break;
282 }
283 if ll < prev_ll - config.tol {
284 }
286 prev_ll = ll;
287 }
288
289 Ok(GmmFitResult {
290 log_likelihood: prev_ll,
291 iterations,
292 converged,
293 })
294 }
295}
296
297impl Next<&[f64]> for GMM {
298 type Output = MarketRegime;
299
300 fn next(&mut self, x: &[f64]) -> Self::Output {
301 let mut max_prob = -1.0;
302 let mut best_k = 0;
303
304 for k in 0..self.k {
305 let p = self.weights[k] * self.pdf(x, k);
306 if p > max_prob {
307 max_prob = p;
308 best_k = k;
309 }
310 }
311
312 match best_k {
313 0 => MarketRegime::Steady,
314 k if k == self.k - 1 => MarketRegime::Crisis,
315 _ => MarketRegime::Cluster(best_k as u8),
316 }
317 }
318}
319
320#[cfg(test)]
321mod tests {
322 use super::*;
323 use approx::assert_relative_eq;
324
325 fn sample_three_gaussians(seed: u64) -> (Vec<Vec<f64>>, Vec<f64>) {
326 let mut data = Vec::new();
327 let true_means = [-5.0, 0.0, 5.0];
328 let mut state = seed;
329 for (c, &mu) in true_means.iter().enumerate() {
330 for _ in 0..200 {
331 state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
332 let u = (state >> 11) as f64 / (1u64 << 53) as f64;
333 let v = (state >> 17) as f64 / (1u64 << 47) as f64;
334 let z = (-2.0 * u.ln()).sqrt() * (2.0 * std::f64::consts::PI * v).cos();
335 data.push(vec![mu + z * 0.5]);
336 let _ = c;
337 }
338 }
339 (data, true_means.to_vec())
340 }
341
342 #[test]
343 fn fit_recovers_three_gaussian_means() {
344 let (data, true_means) = sample_three_gaussians(99);
345 let mut gmm = GMM::with_components(3, 1);
346 let result = gmm
347 .fit(&data, &GmmFitConfig::default())
348 .expect("fit should succeed");
349 assert!(result.converged);
350 let mut recovered: Vec<f64> = gmm.means().iter().map(|m| m[0]).collect();
351 recovered.sort_by(|a, b| a.partial_cmp(b).unwrap());
352 let mut expected = true_means;
353 expected.sort_by(|a, b| a.partial_cmp(b).unwrap());
354 for (r, e) in recovered.iter().zip(expected.iter()) {
355 assert_relative_eq!(r, e, epsilon = 0.75);
356 }
357 for w in gmm.weights() {
358 assert_relative_eq!(*w, 1.0 / 3.0, epsilon = 0.15);
359 }
360 }
361
362 #[test]
363 fn fit_insufficient_data_errors() {
364 let mut gmm = GMM::with_components(3, 1);
365 let err = gmm.fit(&[vec![1.0], vec![2.0]], &GmmFitConfig::default());
366 assert!(matches!(err, Err(GmmError::InsufficientData { .. })));
367 }
368
369 #[test]
370 fn log_likelihood_non_decreasing_on_easy_data() {
371 let (data, _) = sample_three_gaussians(7);
372 let mut gmm = GMM::with_components(3, 1);
373 gmm.validate_data(&data).unwrap();
374 gmm.init_from_quantiles(&data);
375 let mut prev = f64::NEG_INFINITY;
376 for _ in 0..10 {
377 let resp = gmm.responsibilities(&data);
378 gmm.m_step(&data, &resp);
379 let ll = gmm.log_likelihood(&data);
380 assert!(ll >= prev - 1e-9, "LL decreased: {prev} -> {ll}");
381 prev = ll;
382 }
383 }
384
385 #[test]
386 fn max_iter_one_reports_not_converged() {
387 let (data, _) = sample_three_gaussians(3);
388 let mut gmm = GMM::with_components(3, 1);
389 let cfg = GmmFitConfig {
390 max_iter: 1,
391 tol: 1e-12,
392 seed: 1,
393 };
394 let result = gmm.fit(&data, &cfg).unwrap();
395 assert!(!result.converged);
396 assert_eq!(result.iterations, 1);
397 }
398}