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 = 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
120pub 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, ¶ms.means);
142 let vol = mixture_vol(&probs, ¶ms.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
153pub 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 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 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 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#[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
322use 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};