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 =
111        Normal::new(0.0, 1.0).map_err(|e| GaussianHmmError::InvalidParams(e.to_string()))?;
112    let mut out = Vec::with_capacity(n);
113    for t in 0..n {
114        let probs: Vec<f64> = (0..m).map(|s| forward_filter[s][t]).collect();
115        let u = mixture_cdf(observations[t], &probs, params).clamp(1e-12, 1.0 - 1e-12);
116        out.push(normal.inverse_cdf(u));
117    }
118    Ok(out)
119}
120
121/// Weighted mean/vol/lambda per bar from state probabilities (ldhmm `decode_stats_history`).
122pub fn decode_stats_history(
123    params: &GaussianHmmParams,
124    state_probs: &[Vec<f64>],
125) -> Result<Vec<HmmDecodeStatsRow>, GaussianHmmError> {
126    params.validate()?;
127    let m = params.n_states;
128    if state_probs.len() != m {
129        return Err(GaussianHmmError::InvalidParams(
130            "state_probs rows must match n_states".into(),
131        ));
132    }
133    let n = state_probs[0].len();
134    if state_probs.iter().any(|row| row.len() != n) {
135        return Err(GaussianHmmError::InvalidParams(
136            "state_probs rows must have equal length".into(),
137        ));
138    }
139    let mut rows = Vec::with_capacity(n);
140    for t in 0..n {
141        let probs: Vec<f64> = (0..m).map(|s| state_probs[s][t]).collect();
142        let mean = mixture_mean(&probs, &params.means);
143        let vol = mixture_vol(&probs, &params.means, params);
144        let lambda = mixture_lambda(&probs, params);
145        rows.push(HmmDecodeStatsRow {
146            weighted_mean: mean,
147            weighted_vol: vol,
148            weighted_lambda: lambda,
149        });
150    }
151    Ok(rows)
152}
153
154/// Per-state weighted summary stats from observations and state probabilities.
155pub fn calc_stats_from_obs(
156    params: &GaussianHmmParams,
157    observations: &[f64],
158    state_probs: &[Vec<f64>],
159) -> Result<Vec<HmmStateObsStats>, GaussianHmmError> {
160    params.validate()?;
161    let m = params.n_states;
162    let n = observations.len();
163    if state_probs.len() != m || state_probs.first().map(|r| r.len()).unwrap_or(0) != n {
164        return Err(GaussianHmmError::InvalidParams(
165            "state_probs shape must be [n_states][n_obs]".into(),
166        ));
167    }
168    let mut out = Vec::with_capacity(m);
169    for s in 0..m {
170        let mut w_sum = 0.0;
171        let mut mean_acc = 0.0;
172        for t in 0..n {
173            let w = state_probs[s][t];
174            w_sum += w;
175            mean_acc += w * observations[t];
176        }
177        let mean = if w_sum > 0.0 { mean_acc / w_sum } else { 0.0 };
178        let mut var_acc = 0.0;
179        for t in 0..n {
180            let w = state_probs[s][t];
181            var_acc += w * (observations[t] - mean).powi(2);
182        }
183        let vol = if w_sum > 0.0 {
184            (var_acc / w_sum).sqrt()
185        } else {
186            0.0
187        };
188        out.push(HmmStateObsStats {
189            state: s,
190            weight_sum: w_sum,
191            mean,
192            vol,
193            lambda: params.lambdas.get(s).copied().unwrap_or(1.0),
194        });
195    }
196    Ok(out)
197}
198
199impl GaussianHmmParams {
200    /// Forecast state probabilities `h` steps ahead from `current_probs`.
201    pub fn forecast_state(
202        &self,
203        current_probs: &[f64],
204        horizon: usize,
205    ) -> Result<Vec<f64>, GaussianHmmError> {
206        forecast_state(self, current_probs, horizon)
207    }
208
209    /// Forecast mixture volatility `h` steps ahead.
210    pub fn forecast_volatility(
211        &self,
212        current_probs: &[f64],
213        horizon: usize,
214    ) -> Result<f64, GaussianHmmError> {
215        forecast_volatility(self, current_probs, horizon)
216    }
217
218    /// Full forecast/diagnostic bundle after decode.
219    pub fn diagnostics(
220        &self,
221        decode: &GaussianHmmDecode,
222        observations: &[f64],
223    ) -> Result<HmmDiagnostics, GaussianHmmError> {
224        if observations.len() != decode.forward_filter[0].len() {
225            return Err(GaussianHmmError::InvalidParams(
226                "observations length must match decode".into(),
227            ));
228        }
229        let m = self.n_states;
230        let last_probs: Vec<f64> = (0..m)
231            .map(|s| decode.forward_filter[s].last().copied().unwrap_or(0.0))
232            .collect();
233        Ok(HmmDiagnostics {
234            pseudo_residuals: pseudo_residuals(self, &decode.forward_filter, observations)?,
235            decode_stats: decode_stats_history(self, &decode.smooth_probs)?,
236            state_obs_stats: calc_stats_from_obs(self, observations, &decode.smooth_probs)?,
237            forecast_state_h1: forecast_state(self, &last_probs, 1)?,
238            forecast_state_h2: forecast_state(self, &last_probs, 2)?,
239            forecast_vol_h1: forecast_volatility(self, &last_probs, 1)?,
240            forecast_vol_h2: forecast_volatility(self, &last_probs, 2)?,
241            forecast_mean_h1: forecast_observation_mean(self, &last_probs, 1)?,
242        })
243    }
244}
245
246/// Batch diagnostics output.
247#[derive(Debug, Clone, PartialEq)]
248pub struct HmmDiagnostics {
249    pub pseudo_residuals: Vec<f64>,
250    pub decode_stats: Vec<HmmDecodeStatsRow>,
251    pub state_obs_stats: Vec<HmmStateObsStats>,
252    pub forecast_state_h1: Vec<f64>,
253    pub forecast_state_h2: Vec<f64>,
254    pub forecast_vol_h1: f64,
255    pub forecast_vol_h2: f64,
256    pub forecast_mean_h1: f64,
257}
258
259fn transition_power(gamma: &[Vec<f64>], horizon: usize) -> Result<DMatrix<f64>, GaussianHmmError> {
260    let m = gamma.len();
261    if m == 0 {
262        return Err(GaussianHmmError::InvalidParams(
263            "empty transition matrix".into(),
264        ));
265    }
266    let flat: Vec<f64> = gamma.iter().flat_map(|row| row.iter().copied()).collect();
267    if flat.len() != m * m {
268        return Err(GaussianHmmError::InvalidParams(
269            "transition matrix must be square".into(),
270        ));
271    }
272    Ok(DMatrix::from_row_slice(m, m, &flat).pow(horizon as u32))
273}
274
275fn mixture_mean(probs: &[f64], means: &[f64]) -> f64 {
276    probs.iter().zip(means.iter()).map(|(&p, &mu)| p * mu).sum()
277}
278
279fn mixture_vol(probs: &[f64], means: &[f64], params: &GaussianHmmParams) -> f64 {
280    let mean = mixture_mean(probs, means);
281    let var: f64 = probs
282        .iter()
283        .enumerate()
284        .map(|(s, &p)| {
285            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
286            let var_s = ecld_variance(params.stds[s], lam);
287            p * (var_s + (params.means[s] - mean).powi(2))
288        })
289        .sum();
290    var.max(0.0).sqrt()
291}
292
293fn mixture_lambda(probs: &[f64], params: &GaussianHmmParams) -> f64 {
294    probs
295        .iter()
296        .enumerate()
297        .map(|(s, &p)| p * params.lambdas.get(s).copied().unwrap_or(1.0))
298        .sum()
299}
300
301fn mixture_pdf(x: f64, probs: &[f64], params: &GaussianHmmParams) -> f64 {
302    probs
303        .iter()
304        .enumerate()
305        .map(|(s, &p)| {
306            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
307            p * ecld_pdf(x, params.means[s], params.stds[s], lam)
308        })
309        .sum()
310}
311
312fn mixture_cdf(x: f64, probs: &[f64], params: &GaussianHmmParams) -> f64 {
313    probs
314        .iter()
315        .enumerate()
316        .map(|(s, &p)| {
317            let lam = params.lambdas.get(s).copied().unwrap_or(1.0);
318            p * ecld_cdf(x, params.means[s], params.stds[s], lam)
319        })
320        .sum()
321}
322
323// --- IndicatorMetadata (quantwave-i9dn) ---
324
325use crate::indicators::metadata::{IndicatorMetadata, ParamDef};
326
327pub const HMM_FORECAST_METADATA: IndicatorMetadata = IndicatorMetadata {
328    name: "hmm_forecast",
329    description: "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: &[ParamDef {
344        name: "horizon",
345        default: "1",
346        description: "Forecast horizon h (bars ahead).",
347    }],
348    formula_source: "references/ldhmm/ldhmm-cran-reference.pdf; references/ldhmm/ssrn-2979516.pdf",
349    formula_latex: "",
350    gold_standard_file: "hmm_lambda_2state.json",
351    category: "Regime",
352};