Skip to main content

quantwave_core/regimes/
gmm.rs

1//! Gaussian Mixture Models (Two Sigma 2021)
2//!
3//! Source: Two Sigma (2021). "A Machine Learning Approach to Regime Modeling."
4//! Foundational EM Algorithm: Dempster, A. P., Laird, N. M., & Rubin, D. B. (1977).
5//! "Maximum Likelihood from Incomplete Data via the EM Algorithm."
6//! Journal of the Royal Statistical Society: Series B (Methodological), 39(1), 1-22.
7//!
8//! Multi-variate clustering for latent market states using the Expectation-Maximization (EM)
9//! algorithm. This implementation uses diagonal covariance matrices for efficiency.
10
11use 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/// A Gaussian Mixture Model for multi-factor regime detection.
19#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct GMM {
21    k: usize,
22    dims: usize,
23    /// Means for each component [k][dim]
24    means: Vec<Vec<f64>>,
25    /// Variances for each component [k][dim] (diagonal covariance)
26    vars: Vec<Vec<f64>>,
27    /// Mixing coefficients
28    weights: Vec<f64>,
29}
30
31/// Configuration for EM fitting.
32#[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/// Result of EM parameter estimation.
50#[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    /// Creates a new GMM with pre-defined parameters.
69    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    /// Unfitted model with `k` components and `dims` dimensions (for `.fit()`).
82    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    /// Log PDF of x under component `k_idx` (diagonal Gaussian).
112    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    /// Calculate multivariate Gaussian PDF (diagonal covariance)
124    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    /// Batch fit using EM (diagonal covariance).
266    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                // numerical wobble — still accept if close
290            }
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}