Skip to main content

quantwave_core/regimes/
hmm_forecast.rs

1//! HMM forecasting and diagnostics (ldhmm parity — slice 3).
2//!
3//! Implements `forecast_state`, `forecast_volatility`, `forecast_prob`, `pseudo_residuals`,
4//! and decode statistics on top of fitted Gaussian/lambda HMM parameters.
5//!
6//! Sources: references/ldhmm/ldhmm-cran-reference.pdf; references/ldhmm/ssrn-2979516.pdf
7
8use super::ecld::{ecld_cdf, ecld_pdf, ecld_variance};
9use super::gaussian_hmm::{GaussianHmmDecode, GaussianHmmError, GaussianHmmParams};
10use nalgebra::DMatrix;
11use serde::{Deserialize, Serialize};
12use statrs::distribution::{ContinuousCDF, Normal};
13
14/// Per-bar weighted decode statistics (ldhmm `decode_stats_history`).
15#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
16pub struct HmmDecodeStatsRow {
17    pub weighted_mean: f64,
18    pub weighted_vol: f64,
19    pub weighted_lambda: f64,
20}
21
22/// Per-state summary statistics from weighted observations (ldhmm `calc_stats_from_obs`).
23#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
24pub struct HmmStateObsStats {
25    pub state: usize,
26    pub weight_sum: f64,
27    pub mean: f64,
28    pub vol: f64,
29    pub lambda: f64,
30}
31
32/// h-step ahead state probability forecast from a current distribution (ldhmm `forecast_state`).
33///
34/// Returns `π_{t+h|t} = π_t · Γ^h` as a row vector over states.
35pub fn forecast_state(
36    params: &GaussianHmmParams,
37    current_probs: &[f64],
38    horizon: usize,
39) -> Result<Vec<f64>, GaussianHmmError> {
40    params.validate()?;
41    let m = params.n_states;
42    if current_probs.len() != m {
43        return Err(GaussianHmmError::InvalidParams(
44            "current_probs length must match n_states".into(),
45        ));
46    }
47    if horizon == 0 {
48        return Ok(current_probs.to_vec());
49    }
50    let gamma_h = transition_power(&params.gamma, horizon)?;
51    let pi = DMatrix::from_row_slice(1, m, current_probs);
52    let forecast = &pi * gamma_h;
53    Ok(forecast.as_slice().to_vec())
54}
55
56/// Mixture volatility forecast h steps ahead (SSRN 2979516 / ldhmm `forecast_volatility`).
57///
58/// Returns `sqrt(Σ_s π_{t+h|t}(s) · (var_s + (μ_s − μ̄)²))` using lambda-aware emission variances.
59pub fn forecast_volatility(
60    params: &GaussianHmmParams,
61    current_probs: &[f64],
62    horizon: usize,
63) -> Result<f64, GaussianHmmError> {
64    let probs = forecast_state(params, current_probs, horizon)?;
65    Ok(mixture_vol(&probs, &params.means, params))
66}
67
68/// Mixture mean forecast h steps ahead.
69pub fn forecast_observation_mean(
70    params: &GaussianHmmParams,
71    current_probs: &[f64],
72    horizon: usize,
73) -> Result<f64, GaussianHmmError> {
74    let probs = forecast_state(params, current_probs, horizon)?;
75    Ok(mixture_mean(&probs, &params.means))
76}
77
78/// Mixture observation density at `x` for horizon `h` (ldhmm `forecast_prob` point evaluation).
79pub fn forecast_observation_pdf(
80    params: &GaussianHmmParams,
81    current_probs: &[f64],
82    horizon: usize,
83    x: f64,
84) -> Result<f64, GaussianHmmError> {
85    let probs = forecast_state(params, current_probs, horizon)?;
86    Ok(mixture_pdf(x, &probs, params))
87}
88
89/// Probability integral transform residuals (ldhmm `pseudo_residuals`).
90///
91/// Uses the filtered mixture CDF at each time: `u_t = F_t(x_t)`, then `Φ^{-1}(u_t)`.
92pub fn pseudo_residuals(
93    params: &GaussianHmmParams,
94    forward_filter: &[Vec<f64>],
95    observations: &[f64],
96) -> Result<Vec<f64>, GaussianHmmError> {
97    params.validate()?;
98    let m = params.n_states;
99    let n = observations.len();
100    if forward_filter.len() != m {
101        return Err(GaussianHmmError::InvalidParams(
102            "forward_filter rows must match n_states".into(),
103        ));
104    }
105    if forward_filter.first().map(|r| r.len()).unwrap_or(0) != n {
106        return Err(GaussianHmmError::InvalidParams(
107            "forward_filter length must match observations".into(),
108        ));
109    }
110    let normal = Normal::new(0.0, 1.0).expect("standard normal");
111    let mut out = Vec::with_capacity(n);
112    for t in 0..n {
113        let probs: Vec<f64> = (0..m).map(|s| forward_filter[s][t]).collect();
114        let u = mixture_cdf(observations[t], &probs, params).clamp(1e-12, 1.0 - 1e-12);
115        out.push(normal.inverse_cdf(u));
116    }
117    Ok(out)
118}
119
120/// Weighted mean/vol/lambda per bar from state probabilities (ldhmm `decode_stats_history`).
121pub fn decode_stats_history(
122    params: &GaussianHmmParams,
123    state_probs: &[Vec<f64>],
124) -> Result<Vec<HmmDecodeStatsRow>, GaussianHmmError> {
125    params.validate()?;
126    let m = params.n_states;
127    if state_probs.len() != m {
128        return Err(GaussianHmmError::InvalidParams(
129            "state_probs rows must match n_states".into(),
130        ));
131    }
132    let n = state_probs[0].len();
133    if state_probs.iter().any(|row| row.len() != n) {
134        return Err(GaussianHmmError::InvalidParams(
135            "state_probs rows must have equal length".into(),
136        ));
137    }
138    let mut rows = Vec::with_capacity(n);
139    for t in 0..n {
140        let probs: Vec<f64> = (0..m).map(|s| state_probs[s][t]).collect();
141        let mean = mixture_mean(&probs, &params.means);
142        let vol = mixture_vol(&probs, &params.means, params);
143        let lambda = mixture_lambda(&probs, params);
144        rows.push(HmmDecodeStatsRow {
145            weighted_mean: mean,
146            weighted_vol: vol,
147            weighted_lambda: lambda,
148        });
149    }
150    Ok(rows)
151}
152
153/// Per-state weighted summary stats from observations and state probabilities.
154pub fn calc_stats_from_obs(
155    params: &GaussianHmmParams,
156    observations: &[f64],
157    state_probs: &[Vec<f64>],
158) -> Result<Vec<HmmStateObsStats>, GaussianHmmError> {
159    params.validate()?;
160    let m = params.n_states;
161    let n = observations.len();
162    if state_probs.len() != m || state_probs.first().map(|r| r.len()).unwrap_or(0) != n {
163        return Err(GaussianHmmError::InvalidParams(
164            "state_probs shape must be [n_states][n_obs]".into(),
165        ));
166    }
167    let mut out = Vec::with_capacity(m);
168    for s in 0..m {
169        let mut w_sum = 0.0;
170        let mut mean_acc = 0.0;
171        for t in 0..n {
172            let w = state_probs[s][t];
173            w_sum += w;
174            mean_acc += w * observations[t];
175        }
176        let mean = if w_sum > 0.0 { mean_acc / w_sum } else { 0.0 };
177        let mut var_acc = 0.0;
178        for t in 0..n {
179            let w = state_probs[s][t];
180            var_acc += w * (observations[t] - mean).powi(2);
181        }
182        let vol = if w_sum > 0.0 {
183            (var_acc / w_sum).sqrt()
184        } else {
185            0.0
186        };
187        out.push(HmmStateObsStats {
188            state: s,
189            weight_sum: w_sum,
190            mean,
191            vol,
192            lambda: params.lambdas.get(s).copied().unwrap_or(1.0),
193        });
194    }
195    Ok(out)
196}
197
198impl GaussianHmmParams {
199    /// Forecast state probabilities `h` steps ahead from `current_probs`.
200    pub fn forecast_state(
201        &self,
202        current_probs: &[f64],
203        horizon: usize,
204    ) -> Result<Vec<f64>, GaussianHmmError> {
205        forecast_state(self, current_probs, horizon)
206    }
207
208    /// Forecast mixture volatility `h` steps ahead.
209    pub fn forecast_volatility(
210        &self,
211        current_probs: &[f64],
212        horizon: usize,
213    ) -> Result<f64, GaussianHmmError> {
214        forecast_volatility(self, current_probs, horizon)
215    }
216
217    /// Full forecast/diagnostic bundle after decode.
218    pub fn diagnostics(
219        &self,
220        decode: &GaussianHmmDecode,
221        observations: &[f64],
222    ) -> Result<HmmDiagnostics, GaussianHmmError> {
223        if observations.len() != decode.forward_filter[0].len() {
224            return Err(GaussianHmmError::InvalidParams(
225                "observations length must match decode".into(),
226            ));
227        }
228        let m = self.n_states;
229        let last_probs: Vec<f64> = (0..m).map(|s| decode.forward_filter[s].last().copied().unwrap_or(0.0)).collect();
230        Ok(HmmDiagnostics {
231            pseudo_residuals: pseudo_residuals(self, &decode.forward_filter, observations)?,
232            decode_stats: decode_stats_history(self, &decode.smooth_probs)?,
233            state_obs_stats: calc_stats_from_obs(self, observations, &decode.smooth_probs)?,
234            forecast_state_h1: forecast_state(self, &last_probs, 1)?,
235            forecast_state_h2: forecast_state(self, &last_probs, 2)?,
236            forecast_vol_h1: forecast_volatility(self, &last_probs, 1)?,
237            forecast_vol_h2: forecast_volatility(self, &last_probs, 2)?,
238            forecast_mean_h1: forecast_observation_mean(self, &last_probs, 1)?,
239        })
240    }
241}
242
243/// Batch diagnostics output.
244#[derive(Debug, Clone, PartialEq)]
245pub struct HmmDiagnostics {
246    pub pseudo_residuals: Vec<f64>,
247    pub decode_stats: Vec<HmmDecodeStatsRow>,
248    pub state_obs_stats: Vec<HmmStateObsStats>,
249    pub forecast_state_h1: Vec<f64>,
250    pub forecast_state_h2: Vec<f64>,
251    pub forecast_vol_h1: f64,
252    pub forecast_vol_h2: f64,
253    pub forecast_mean_h1: f64,
254}
255
256fn transition_power(gamma: &[Vec<f64>], horizon: usize) -> Result<DMatrix<f64>, GaussianHmmError> {
257    let m = gamma.len();
258    if m == 0 {
259        return Err(GaussianHmmError::InvalidParams("empty transition matrix".into()));
260    }
261    let flat: Vec<f64> = gamma.iter().flat_map(|row| row.iter().copied()).collect();
262    if flat.len() != m * m {
263        return Err(GaussianHmmError::InvalidParams(
264            "transition matrix must be square".into(),
265        ));
266    }
267    Ok(DMatrix::from_row_slice(m, m, &flat).pow(horizon as u32))
268}
269
270fn mixture_mean(probs: &[f64], means: &[f64]) -> f64 {
271    probs
272        .iter()
273        .zip(means.iter())
274        .map(|(&p, &mu)| p * mu)
275        .sum()
276}
277
278fn mixture_vol(probs: &[f64], means: &[f64], params: &GaussianHmmParams) -> f64 {
279    let mean = mixture_mean(probs, means);
280    let var: f64 = probs
281        .iter()
282        .enumerate()
283        .map(|(s, &p)| {
284            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
285            let var_s = ecld_variance(params.stds[s], lam);
286            p * (var_s + (params.means[s] - mean).powi(2))
287        })
288        .sum();
289    var.max(0.0).sqrt()
290}
291
292fn mixture_lambda(probs: &[f64], params: &GaussianHmmParams) -> f64 {
293    probs
294        .iter()
295        .enumerate()
296        .map(|(s, &p)| p * params.lambdas.get(s).copied().unwrap_or(1.0))
297        .sum()
298}
299
300fn mixture_pdf(x: f64, probs: &[f64], params: &GaussianHmmParams) -> f64 {
301    probs
302        .iter()
303        .enumerate()
304        .map(|(s, &p)| {
305            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
306            p * ecld_pdf(x, params.means[s], params.stds[s], lam)
307        })
308        .sum()
309}
310
311fn mixture_cdf(x: f64, probs: &[f64], params: &GaussianHmmParams) -> f64 {
312    probs
313        .iter()
314        .enumerate()
315        .map(|(s, &p)| {
316            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
317            p * ecld_cdf(x, params.means[s], params.stds[s], lam)
318        })
319        .sum()
320}
321
322// --- IndicatorMetadata (quantwave-i9dn) ---
323
324use crate::indicators::metadata::{IndicatorMetadata, ParamDef};
325
326pub const HMM_FORECAST_METADATA: IndicatorMetadata = IndicatorMetadata {
327    name: "hmm_forecast",
328    description:
329        "HMM forecasting and diagnostics: state/vol/probability forecasts, pseudo-residuals, decode stats.",
330    usage: "After fitting/decoding a Gaussian or lambda HMM, forecast regimes and volatility h steps ahead, \
331             evaluate mixture predictive densities, and extract pseudo-residuals for model checking.",
332    keywords: &[
333        "regime",
334        "hmm",
335        "forecast",
336        "volatility",
337        "pseudo_residuals",
338        "ldhmm",
339        "diagnostics",
340    ],
341    ehlers_summary: "ldhmm-style post-fit analytics on homogeneous HMMs: π_{t+h|t}=π_t·Γ^h state forecasts, \
342                     mixture volatility per SSRN 2979516, filtered pseudo-residuals, and decode_stats_history.",
343    params: &[
344        ParamDef {
345            name: "horizon",
346            default: "1",
347            description: "Forecast horizon h (bars ahead).",
348        },
349    ],
350    formula_source: "references/ldhmm/ldhmm-cran-reference.pdf; references/ldhmm/ssrn-2979516.pdf",
351    formula_latex: "",
352    gold_standard_file: "hmm_lambda_2state.json",
353    category: "Regime",
354};