use serde::Serialize;
use crate::stats::{mean, std_dev, variance};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RegimeCompareOpts {
pub zero_tol: f64,
pub min_periods: usize,
pub tie_tol: f64,
}
impl Default for RegimeCompareOpts {
fn default() -> Self {
Self {
zero_tol: 1e-9,
min_periods: 8,
tie_tol: 1e-12,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Serialize)]
pub struct ZagaSplit {
pub n: usize,
pub zero_mass: f64,
pub n_nonzero: usize,
pub positive_share: f64,
pub cont_mean: f64,
pub cont_sd: f64,
pub cont_median: f64,
pub gamma_shape: f64,
pub gamma_rate: f64,
pub pooled_mean: f64,
}
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct RegimeComparison {
pub regime: String,
pub n_periods: usize,
pub a: ZagaSplit,
pub b: ZagaSplit,
pub zero_mass_gap: f64,
pub mean_gap: f64,
pub cont_mean_gap: f64,
pub ks_statistic: f64,
pub edge_sign: i8,
pub counted: bool,
}
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct RegimeDistributionReport {
pub regimes: Vec<RegimeComparison>,
pub pooled_mean_gap: f64,
pub pooled_edge_sign: i8,
pub reversal_regimes: Vec<String>,
pub pooled_hides_reversal: bool,
pub edge_dispersion: f64,
}
fn sign_of(x: f64, tie_tol: f64) -> i8 {
if x > tie_tol {
1
} else if x < -tie_tol {
-1
} else {
0
}
}
fn sorted_copy(xs: &[f64]) -> Vec<f64> {
let mut v = xs.to_vec();
v.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
v
}
fn median_of(xs: &[f64]) -> f64 {
let n = xs.len();
if n == 0 {
return 0.0;
}
let s = sorted_copy(xs);
if n % 2 == 1 {
s[n / 2]
} else {
0.5 * (s[n / 2 - 1] + s[n / 2])
}
}
fn ks_two_sample(a: &[f64], b: &[f64]) -> f64 {
if a.is_empty() || b.is_empty() {
return 0.0;
}
let xa = sorted_copy(a);
let xb = sorted_copy(b);
let na = xa.len() as f64;
let nb = xb.len() as f64;
let (mut i, mut j) = (0usize, 0usize);
let mut d = 0.0_f64;
loop {
let v = match (xa.get(i), xb.get(j)) {
(Some(&va), Some(&vb)) => va.min(vb),
(Some(&va), None) => va,
(None, Some(&vb)) => vb,
(None, None) => break,
};
while i < xa.len() && xa[i] <= v {
i += 1;
}
while j < xb.len() && xb[j] <= v {
j += 1;
}
let diff = (i as f64 / na - j as f64 / nb).abs();
if diff > d {
d = diff;
}
}
d
}
fn zaga_split(xs: &[f64], zero_tol: f64) -> ZagaSplit {
let n = xs.len();
let cont: Vec<f64> = xs.iter().copied().filter(|r| r.abs() > zero_tol).collect();
let n_nonzero = cont.len();
let zero_mass = if n == 0 {
0.0
} else {
(n - n_nonzero) as f64 / n as f64
};
let positive_share = if n_nonzero == 0 {
0.0
} else {
cont.iter().filter(|&&r| r > 0.0).count() as f64 / n_nonzero as f64
};
let mags: Vec<f64> = cont.iter().map(|r| r.abs()).collect();
let mm = mean(&mags);
let mv = variance(&mags);
let (gamma_shape, gamma_rate) = if mm > 0.0 && mv > 0.0 {
(mm * mm / mv, mm / mv)
} else {
(0.0, 0.0)
};
ZagaSplit {
n,
zero_mass,
n_nonzero,
positive_share,
cont_mean: mean(&cont),
cont_sd: std_dev(&cont),
cont_median: median_of(&cont),
gamma_shape,
gamma_rate,
pooled_mean: mean(xs),
}
}
pub fn compare_by_regime(
a: &[f64],
b: &[f64],
regimes: &[&str],
opts: RegimeCompareOpts,
) -> RegimeDistributionReport {
let n = a.len().min(b.len()).min(regimes.len());
if n == 0 {
return RegimeDistributionReport {
regimes: Vec::new(),
pooled_mean_gap: 0.0,
pooled_edge_sign: 0,
reversal_regimes: Vec::new(),
pooled_hides_reversal: false,
edge_dispersion: 0.0,
};
}
let mut labels: Vec<String> = regimes[..n].iter().map(|s| (*s).to_string()).collect();
labels.sort();
labels.dedup();
let pooled_mean_gap = mean(&a[..n]) - mean(&b[..n]);
let pooled_edge_sign = sign_of(pooled_mean_gap, opts.tie_tol);
let mut out: Vec<RegimeComparison> = Vec::with_capacity(labels.len());
for label in labels {
let idx: Vec<usize> = (0..n).filter(|&i| regimes[i] == label).collect();
let ra: Vec<f64> = idx.iter().map(|&i| a[i]).collect();
let rb: Vec<f64> = idx.iter().map(|&i| b[i]).collect();
let sa = zaga_split(&ra, opts.zero_tol);
let sb = zaga_split(&rb, opts.zero_tol);
let mean_gap = sa.pooled_mean - sb.pooled_mean;
let cont_a: Vec<f64> = ra
.iter()
.copied()
.filter(|r| r.abs() > opts.zero_tol)
.collect();
let cont_b: Vec<f64> = rb
.iter()
.copied()
.filter(|r| r.abs() > opts.zero_tol)
.collect();
out.push(RegimeComparison {
regime: label,
n_periods: idx.len(),
zero_mass_gap: sa.zero_mass - sb.zero_mass,
mean_gap,
cont_mean_gap: sa.cont_mean - sb.cont_mean,
ks_statistic: ks_two_sample(&cont_a, &cont_b),
edge_sign: sign_of(mean_gap, opts.tie_tol),
counted: idx.len() >= opts.min_periods,
a: sa,
b: sb,
});
}
let reversal_regimes: Vec<String> = out
.iter()
.filter(|r| {
r.counted
&& r.edge_sign != 0
&& pooled_edge_sign != 0
&& r.edge_sign != pooled_edge_sign
})
.map(|r| r.regime.clone())
.collect();
let counted_gaps: Vec<f64> = out
.iter()
.filter(|r| r.counted)
.map(|r| r.mean_gap)
.collect();
let edge_dispersion = if counted_gaps.len() < 2 {
0.0
} else {
let hi = counted_gaps
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
let lo = counted_gaps.iter().copied().fold(f64::INFINITY, f64::min);
hi - lo
};
RegimeDistributionReport {
regimes: out,
pooled_mean_gap,
pooled_edge_sign,
pooled_hides_reversal: !reversal_regimes.is_empty(),
reversal_regimes,
edge_dispersion,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn labels(n: usize) -> Vec<&'static str> {
(0..n)
.map(|i| if i % 2 == 0 { "bull" } else { "bear" })
.collect()
}
#[test]
fn pooling_hides_a_sign_reversal() {
let n = 120;
let lab = labels(n);
let a: Vec<f64> = (0..n)
.map(|i| if i % 2 == 0 { 0.010 } else { -0.010 })
.collect();
let b: Vec<f64> = (0..n)
.map(|i| if i % 2 == 0 { -0.010 } else { 0.010 })
.collect();
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
assert!(
rep.pooled_mean_gap.abs() < 1e-12,
"pooled gap should vanish: {}",
rep.pooled_mean_gap
);
assert_eq!(rep.regimes.len(), 2);
let bull = rep.regimes.iter().find(|r| r.regime == "bull").unwrap();
let bear = rep.regimes.iter().find(|r| r.regime == "bear").unwrap();
assert_eq!(bull.edge_sign, 1);
assert_eq!(bear.edge_sign, -1);
assert!(rep.edge_dispersion > 0.03, "{}", rep.edge_dispersion);
}
#[test]
fn reversal_is_flagged_against_a_signed_pooled_verdict() {
let mut a = Vec::new();
let mut b = Vec::new();
let mut lab: Vec<&'static str> = Vec::new();
for _ in 0..60 {
a.push(0.01);
b.push(0.001);
lab.push("calm");
}
for _ in 0..20 {
a.push(-0.02);
b.push(0.005);
lab.push("crisis");
}
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
assert_eq!(rep.pooled_edge_sign, 1, "A leads pooled");
assert!(rep.pooled_hides_reversal);
assert_eq!(rep.reversal_regimes, vec!["crisis".to_string()]);
}
#[test]
fn consistent_edge_reports_no_reversal() {
let n = 100;
let lab = labels(n);
let a: Vec<f64> = (0..n).map(|i| 0.004 + 0.001 * (i as f64).sin()).collect();
let b: Vec<f64> = (0..n).map(|i| 0.001 + 0.001 * (i as f64).sin()).collect();
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
assert!(!rep.pooled_hides_reversal);
assert!(rep.reversal_regimes.is_empty());
assert!(rep.regimes.iter().all(|r| r.edge_sign == 1));
}
#[test]
fn zero_mass_separates_a_selective_strategy_from_a_timid_one() {
let n = 80;
let lab = vec!["all"; n];
let a: Vec<f64> = (0..n)
.map(|i| if i % 4 == 0 { 0.04 } else { 0.0 })
.collect();
let b: Vec<f64> = vec![0.01; n];
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
let r = &rep.regimes[0];
assert!(
r.mean_gap.abs() < 1e-12,
"pooled means match: {}",
r.mean_gap
);
assert!(
(r.a.zero_mass - 0.75).abs() < 1e-12,
"A sits out 3 of 4: {}",
r.a.zero_mass
);
assert!(r.b.zero_mass.abs() < 1e-12, "B always trades");
assert!(
r.zero_mass_gap > 0.7,
"the split should shout: {}",
r.zero_mass_gap
);
assert!(
(r.cont_mean_gap - 0.03).abs() < 1e-12,
"per-trade gap is 4% vs 1%: {}",
r.cont_mean_gap
);
assert!(
r.ks_statistic > 0.9,
"disjoint supports: {}",
r.ks_statistic
);
}
#[test]
fn ks_catches_a_shape_difference_at_equal_means() {
let n = 100;
let lab = vec!["all"; n];
let a: Vec<f64> = (0..n)
.map(|i| if i % 2 == 0 { 0.001 } else { -0.001 })
.collect();
let b: Vec<f64> = (0..n)
.map(|i| if i % 2 == 0 { 0.05 } else { -0.05 })
.collect();
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
let r = &rep.regimes[0];
assert!(r.mean_gap.abs() < 1e-12);
assert!(
r.ks_statistic > 0.4,
"shape difference should register: {}",
r.ks_statistic
);
assert!(r.b.cont_sd > r.a.cont_sd);
}
#[test]
fn gamma_moments_match_a_known_magnitude_sample() {
let a = vec![1.0, -2.0, 3.0, -4.0];
let b = vec![1.0, -2.0, 3.0, -4.0];
let lab = vec!["all"; 4];
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
let s = rep.regimes[0].a;
assert!((s.gamma_shape - 3.75).abs() < 1e-9, "{}", s.gamma_shape);
assert!((s.gamma_rate - 1.5).abs() < 1e-9, "{}", s.gamma_rate);
assert!((s.positive_share - 0.5).abs() < 1e-12);
}
#[test]
fn short_regimes_do_not_drive_the_reversal_verdict() {
let mut a = vec![0.01; 60];
let mut b = vec![0.001; 60];
let mut lab: Vec<&'static str> = vec!["calm"; 60];
a.extend([-0.05, -0.05, -0.05]);
b.extend([0.02, 0.02, 0.02]);
lab.extend(["crisis", "crisis", "crisis"]);
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
let crisis = rep.regimes.iter().find(|r| r.regime == "crisis").unwrap();
assert_eq!(crisis.edge_sign, -1, "the reversal is still reported");
assert!(!crisis.counted, "but it is not counted");
assert!(!rep.pooled_hides_reversal);
}
#[test]
fn misaligned_inputs_truncate_to_the_shortest() {
let a = vec![0.01; 10];
let b = vec![0.0; 4];
let lab = vec!["all"; 7];
let rep = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
assert_eq!(rep.regimes[0].n_periods, 4);
}
#[test]
fn empty_input_is_inert() {
let rep = compare_by_regime(&[], &[], &[], RegimeCompareOpts::default());
assert!(rep.regimes.is_empty());
assert_eq!(rep.pooled_edge_sign, 0);
assert!(!rep.pooled_hides_reversal);
}
#[test]
fn output_order_is_lexicographic_and_reproducible() {
let lab = vec!["zulu", "alpha", "mike", "alpha", "zulu", "mike"];
let a = vec![0.01, 0.02, 0.03, 0.04, 0.05, 0.06];
let b = vec![0.0; 6];
let first = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
let second = compare_by_regime(&a, &b, &lab, RegimeCompareOpts::default());
assert_eq!(first, second, "recompute must be identical");
let order: Vec<&str> = first.regimes.iter().map(|r| r.regime.as_str()).collect();
assert_eq!(order, vec!["alpha", "mike", "zulu"]);
}
}