1use 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#[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#[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
32pub 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(¶ms.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
56pub 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, ¶ms.means, params))
66}
67
68pub 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, ¶ms.means))
76}
77
78pub 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
89pub 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
121pub 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, ¶ms.means);
143 let vol = mixture_vol(&probs, ¶ms.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
154pub 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 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 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 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#[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
323use 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};