Skip to main content

kestrel_chartkit/
regime_advanced.rs

1//! Advanced regime-model building blocks layered on top of [`crate::regime::classify_regime`]'s
2//! single-shot four-class output: empirical Markov transition probabilities, trend persistence
3//! (streak length), a streaming predictability index, hysteretic (chatter-free) level transitions,
4//! and adaptive swing-length tracking as a dominant-cycle-length proxy.
5
6use std::collections::{HashMap, VecDeque};
7
8use crate::model::MarketRegime;
9use crate::stats::linear_regression;
10
11/// Empirical Markov transition model over [`MarketRegime`] states: learns `P(to | from)` from an
12/// observed regime sequence rather than assuming a fixed transition structure.
13#[derive(Debug, Clone, Default)]
14pub struct RegimeMarkovModel {
15    counts: HashMap<(MarketRegime, MarketRegime), u32>,
16    totals: HashMap<MarketRegime, u32>,
17}
18
19impl RegimeMarkovModel {
20    pub fn new() -> Self {
21        Self::default()
22    }
23
24    /// Records one observed `from -> to` transition (call with `from == to` for "stayed in the
25    /// same regime this bar" too, so self-transition probability is learned as well).
26    pub fn observe_transition(&mut self, from: MarketRegime, to: MarketRegime) {
27        *self.counts.entry((from, to)).or_insert(0) += 1;
28        *self.totals.entry(from).or_insert(0) += 1;
29    }
30
31    /// Empirical `P(to | from)`. `0.0` if `from` has never been observed.
32    pub fn transition_probability(&self, from: MarketRegime, to: MarketRegime) -> f64 {
33        let total = match self.totals.get(&from) {
34            Some(&t) if t > 0 => t,
35            _ => return 0.0,
36        };
37        let count = self.counts.get(&(from, to)).copied().unwrap_or(0);
38        count as f64 / total as f64
39    }
40
41    /// Full distribution over next states given `from`, most probable first. Empty if `from` has
42    /// never been observed.
43    pub fn next_state_distribution(&self, from: MarketRegime) -> Vec<(MarketRegime, f64)> {
44        let states = [
45            MarketRegime::BullishExpansion,
46            MarketRegime::BearishExpansion,
47            MarketRegime::Consolidation,
48            MarketRegime::Transition,
49        ];
50        let mut dist: Vec<(MarketRegime, f64)> = states
51            .into_iter()
52            .map(|to| (to, self.transition_probability(from, to)))
53            .filter(|(_, p)| *p > 0.0)
54            .collect();
55        dist.sort_by(|a, b| b.1.total_cmp(&a.1));
56        dist
57    }
58}
59
60/// Tracks how long the current regime has persisted and, from history, its typical persistence.
61#[derive(Debug, Clone)]
62pub struct RegimePersistenceTracker {
63    current: Option<MarketRegime>,
64    bars_in_regime: u32,
65    completed_streaks: VecDeque<u32>,
66    max_history: usize,
67}
68
69#[derive(Debug, Clone, Copy, PartialEq)]
70pub struct RegimePersistenceOutput {
71    pub bars_in_regime: u32,
72    /// `true` on the bar the regime changed (from the second observed regime onward).
73    pub changed: bool,
74    /// Mean length (in bars) of completed regime streaks so far. `None` until at least one
75    /// regime change has completed a streak.
76    pub average_streak_length: Option<f64>,
77}
78
79impl RegimePersistenceTracker {
80    pub fn new(max_history: usize) -> Self {
81        Self {
82            current: None,
83            bars_in_regime: 0,
84            completed_streaks: VecDeque::new(),
85            max_history: max_history.max(1),
86        }
87    }
88
89    pub fn reset(&mut self) {
90        self.current = None;
91        self.bars_in_regime = 0;
92        self.completed_streaks.clear();
93    }
94
95    pub fn update(&mut self, regime: MarketRegime) -> RegimePersistenceOutput {
96        let changed = match self.current {
97            Some(prev) if prev != regime => {
98                if self.completed_streaks.len() >= self.max_history {
99                    self.completed_streaks.pop_front();
100                }
101                self.completed_streaks.push_back(self.bars_in_regime);
102                self.bars_in_regime = 0;
103                true
104            }
105            None => false,
106            _ => false,
107        };
108
109        self.current = Some(regime);
110        self.bars_in_regime += 1;
111
112        let average_streak_length = if self.completed_streaks.is_empty() {
113            None
114        } else {
115            Some(
116                self.completed_streaks.iter().copied().sum::<u32>() as f64
117                    / self.completed_streaks.len() as f64,
118            )
119        };
120
121        RegimePersistenceOutput {
122            bars_in_regime: self.bars_in_regime,
123            changed,
124            average_streak_length,
125        }
126    }
127}
128
129/// Streaming predictability index: the R^2 of an OLS fit over the trailing `window_len` closes —
130/// close to `1.0` means a clean, linear trend; close to `0.0` means noisy/directionless movement.
131#[derive(Debug, Clone)]
132pub struct PredictabilityTracker {
133    window_len: usize,
134    buffer: VecDeque<f64>,
135}
136
137impl PredictabilityTracker {
138    pub fn new(window_len: usize) -> Self {
139        let window_len = window_len.max(2);
140        Self {
141            window_len,
142            buffer: VecDeque::with_capacity(window_len),
143        }
144    }
145
146    pub fn reset(&mut self) {
147        self.buffer.clear();
148    }
149
150    /// Returns `None` until `window_len` values have been fed.
151    pub fn update(&mut self, close: f64) -> Option<f64> {
152        if self.buffer.len() >= self.window_len {
153            self.buffer.pop_front();
154        }
155        self.buffer.push_back(close);
156        if self.buffer.len() < self.window_len {
157            return None;
158        }
159        let values: Vec<f64> = self.buffer.iter().copied().collect();
160        linear_regression(&values).map(|r| r.r2)
161    }
162}
163
164/// Three-level hysteretic (Schmitt-trigger-style) classification, so a score oscillating near a
165/// single threshold does not flap the classified level back and forth every bar.
166///
167/// Contract: `enter_low <= exit_low <= exit_high <= enter_high`. From [`HysteresisLevel::Neutral`]
168/// a value must reach `enter_high`/`enter_low` to leave; from [`HysteresisLevel::High`]/
169/// [`HysteresisLevel::Low`] a value must retreat past `exit_high`/`exit_low` to fall back to
170/// `Neutral` (an extreme-to-extreme jump takes two updates: extreme -> Neutral -> the other
171/// extreme, never a same-bar jump between the two extremes).
172#[derive(Debug, Clone, Copy, PartialEq, Eq)]
173#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
174pub enum HysteresisLevel {
175    Low,
176    Neutral,
177    High,
178}
179
180#[derive(Debug, Clone, Copy)]
181pub struct HysteresisBand {
182    enter_low: f64,
183    exit_low: f64,
184    exit_high: f64,
185    enter_high: f64,
186    level: HysteresisLevel,
187}
188
189impl HysteresisBand {
190    pub fn new(enter_low: f64, exit_low: f64, exit_high: f64, enter_high: f64) -> Self {
191        Self {
192            enter_low,
193            exit_low,
194            exit_high,
195            enter_high,
196            level: HysteresisLevel::Neutral,
197        }
198    }
199
200    pub fn level(&self) -> HysteresisLevel {
201        self.level
202    }
203
204    pub fn reset(&mut self) {
205        self.level = HysteresisLevel::Neutral;
206    }
207
208    pub fn update(&mut self, value: f64) -> HysteresisLevel {
209        self.level = match self.level {
210            HysteresisLevel::High => {
211                if value <= self.exit_high {
212                    HysteresisLevel::Neutral
213                } else {
214                    HysteresisLevel::High
215                }
216            }
217            HysteresisLevel::Low => {
218                if value >= self.exit_low {
219                    HysteresisLevel::Neutral
220                } else {
221                    HysteresisLevel::Low
222                }
223            }
224            HysteresisLevel::Neutral => {
225                if value >= self.enter_high {
226                    HysteresisLevel::High
227                } else if value <= self.enter_low {
228                    HysteresisLevel::Low
229                } else {
230                    HysteresisLevel::Neutral
231                }
232            }
233        };
234        self.level
235    }
236}
237
238/// Adaptive swing-length tracking: counts bars between sign crossings of a signed input (an
239/// oscillator or signed regime score) as an empirical, adapting proxy for the dominant cycle
240/// length, without claiming spectral/Hilbert-transform precision.
241#[derive(Debug, Clone)]
242pub struct AdaptiveCycleTracker {
243    prev_sign: Option<i8>,
244    bars_since_crossing: u32,
245    recent_swing_lengths: VecDeque<u32>,
246    max_history: usize,
247}
248
249#[derive(Debug, Clone, Copy, PartialEq)]
250pub struct AdaptiveCycleOutput {
251    pub bars_since_crossing: u32,
252    /// Mean bars-between-crossings over recent swings. `None` until at least one full swing
253    /// (crossing to crossing) has been observed.
254    pub average_swing_length: Option<f64>,
255}
256
257impl AdaptiveCycleTracker {
258    pub fn new(max_history: usize) -> Self {
259        Self {
260            prev_sign: None,
261            bars_since_crossing: 0,
262            recent_swing_lengths: VecDeque::new(),
263            max_history: max_history.max(1),
264        }
265    }
266
267    pub fn reset(&mut self) {
268        self.prev_sign = None;
269        self.bars_since_crossing = 0;
270        self.recent_swing_lengths.clear();
271    }
272
273    pub fn update(&mut self, value: f64) -> AdaptiveCycleOutput {
274        let sign: i8 = if value > 0.0 {
275            1
276        } else if value < 0.0 {
277            -1
278        } else {
279            0
280        };
281
282        // A crossing bar is the first bar of the new swing, so the completed swing's length is
283        // `bars_since_crossing` *before* this bar's increment below.
284        let crossed =
285            matches!(self.prev_sign, Some(prev) if prev != 0 && sign != 0 && prev != sign);
286        if crossed {
287            if self.recent_swing_lengths.len() >= self.max_history {
288                self.recent_swing_lengths.pop_front();
289            }
290            self.recent_swing_lengths
291                .push_back(self.bars_since_crossing);
292            self.bars_since_crossing = 0;
293        }
294
295        self.bars_since_crossing += 1;
296
297        if sign != 0 {
298            self.prev_sign = Some(sign);
299        }
300
301        let average_swing_length = if self.recent_swing_lengths.is_empty() {
302            None
303        } else {
304            Some(
305                self.recent_swing_lengths.iter().copied().sum::<u32>() as f64
306                    / self.recent_swing_lengths.len() as f64,
307            )
308        };
309
310        AdaptiveCycleOutput {
311            bars_since_crossing: self.bars_since_crossing,
312            average_swing_length,
313        }
314    }
315}
316
317#[cfg(test)]
318mod tests {
319    use super::*;
320
321    #[test]
322    fn test_markov_model_learns_empirical_transitions() {
323        let mut model = RegimeMarkovModel::new();
324        let sequence = [
325            MarketRegime::BullishExpansion,
326            MarketRegime::BullishExpansion,
327            MarketRegime::Consolidation,
328            MarketRegime::BullishExpansion,
329            MarketRegime::BullishExpansion,
330        ];
331        for pair in sequence.windows(2) {
332            model.observe_transition(pair[0], pair[1]);
333        }
334
335        // BullishExpansion -> BullishExpansion happened twice, -> Consolidation once (3 total).
336        let p_stay = model.transition_probability(
337            MarketRegime::BullishExpansion,
338            MarketRegime::BullishExpansion,
339        );
340        assert!((p_stay - 2.0 / 3.0).abs() < 1e-9);
341
342        let dist = model.next_state_distribution(MarketRegime::BullishExpansion);
343        assert_eq!(dist[0].0, MarketRegime::BullishExpansion);
344
345        // Never-observed source state has an empty distribution.
346        assert!(model
347            .next_state_distribution(MarketRegime::BearishExpansion)
348            .is_empty());
349    }
350
351    #[test]
352    fn test_persistence_tracker_counts_streaks_and_changes() {
353        let mut tracker = RegimePersistenceTracker::new(10);
354        let regimes = [
355            MarketRegime::BullishExpansion,
356            MarketRegime::BullishExpansion,
357            MarketRegime::BullishExpansion,
358            MarketRegime::Consolidation,
359            MarketRegime::Consolidation,
360        ];
361        let mut outputs = Vec::new();
362        for regime in regimes {
363            outputs.push(tracker.update(regime));
364        }
365
366        assert!(!outputs[0].changed);
367        assert_eq!(outputs[2].bars_in_regime, 3);
368        assert!(outputs[3].changed, "regime change must be flagged");
369        assert_eq!(outputs[3].bars_in_regime, 1);
370        assert_eq!(outputs[3].average_streak_length, Some(3.0));
371    }
372
373    #[test]
374    fn test_predictability_tracker_scores_clean_trend_high() {
375        let mut tracker = PredictabilityTracker::new(10);
376        let mut last = None;
377        for i in 0..10 {
378            last = tracker.update(100.0 + i as f64);
379        }
380        assert!(
381            last.unwrap() > 0.99,
382            "a perfectly linear trend must score near 1.0"
383        );
384    }
385
386    #[test]
387    fn test_predictability_tracker_scores_noise_low() {
388        let mut tracker = PredictabilityTracker::new(6);
389        let mut last = None;
390        for v in [100.0, 105.0, 98.0, 107.0, 96.0, 109.0] {
391            last = tracker.update(v);
392        }
393        assert!(
394            last.unwrap() < 0.3,
395            "zig-zagging noise must score low predictability"
396        );
397    }
398
399    #[test]
400    fn test_hysteresis_band_does_not_chatter_near_a_single_threshold() {
401        let mut band = HysteresisBand::new(-2.0, -1.0, 1.0, 2.0);
402        assert_eq!(band.update(2.5), HysteresisLevel::High);
403        // Oscillates between the enter/exit thresholds without ever falling to/below exit_high:
404        // must stay High the whole time (no chatter).
405        for v in [1.8, 1.2, 1.9, 1.3, 1.7] {
406            assert_eq!(band.update(v), HysteresisLevel::High);
407        }
408        // Now genuinely retreats past exit_high -> Neutral.
409        assert_eq!(band.update(0.5), HysteresisLevel::Neutral);
410    }
411
412    #[test]
413    fn test_hysteresis_band_extreme_to_extreme_passes_through_neutral() {
414        let mut band = HysteresisBand::new(-2.0, -1.0, 1.0, 2.0);
415        assert_eq!(band.update(2.5), HysteresisLevel::High);
416        assert_eq!(band.update(-2.5), HysteresisLevel::Neutral);
417        assert_eq!(band.update(-2.5), HysteresisLevel::Low);
418    }
419
420    #[test]
421    fn test_adaptive_cycle_tracker_measures_swing_length() {
422        let mut tracker = AdaptiveCycleTracker::new(10);
423        // +,+,+,+ (4 bars) then -,-,-,- (4 bars): one sign crossing after 4 bars.
424        let values = [1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, -1.0];
425        let mut last = AdaptiveCycleOutput {
426            bars_since_crossing: 0,
427            average_swing_length: None,
428        };
429        for v in values {
430            last = tracker.update(v);
431        }
432        assert_eq!(last.average_swing_length, Some(4.0));
433        assert_eq!(last.bars_since_crossing, 4);
434    }
435}