use crate::error::{ForecastError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PanelAggregator {
Mean,
Median,
Std,
Rank,
}
pub fn panel_aggregate(
values: &[Vec<f64>],
kind: PanelAggregator,
exclude_self: bool,
) -> Result<Vec<Vec<f64>>> {
let n = values.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
let t = values[0].len();
for v in values.iter() {
if v.len() != t {
return Err(ForecastError::DimensionMismatch {
expected: t,
got: v.len(),
});
}
}
if exclude_self && n < 2 {
return Err(ForecastError::InsufficientData {
needed: 2,
got: n,
hint: Some("exclude_self requires ≥ 2 series".into()),
});
}
let result = match kind {
PanelAggregator::Rank => rank_per_series(values, t),
_ => {
if exclude_self {
loo_aggregate(values, kind, n, t)?
} else {
let common = full_aggregate(values, kind, n, t)?;
vec![common; n]
}
}
};
Ok(result)
}
pub fn panel_mean(values: &[Vec<f64>], exclude_self: bool) -> Result<Vec<Vec<f64>>> {
panel_aggregate(values, PanelAggregator::Mean, exclude_self)
}
pub fn panel_median(values: &[Vec<f64>], exclude_self: bool) -> Result<Vec<Vec<f64>>> {
panel_aggregate(values, PanelAggregator::Median, exclude_self)
}
pub fn panel_std(values: &[Vec<f64>], exclude_self: bool) -> Result<Vec<Vec<f64>>> {
panel_aggregate(values, PanelAggregator::Std, exclude_self)
}
pub fn panel_rank(values: &[Vec<f64>]) -> Result<Vec<Vec<f64>>> {
panel_aggregate(values, PanelAggregator::Rank, false)
}
fn full_aggregate(
values: &[Vec<f64>],
kind: PanelAggregator,
n: usize,
t: usize,
) -> Result<Vec<f64>> {
if matches!(kind, PanelAggregator::Std) && n < 2 {
return Err(ForecastError::InsufficientData {
needed: 2,
got: n,
hint: Some("Std needs ≥ 2 series for sample variance".into()),
});
}
let mut out = vec![0.0; t];
let mut col = vec![0.0_f64; n];
for j in 0..t {
for (i, series) in values.iter().enumerate() {
col[i] = series[j];
}
out[j] = compute_kind(&col, kind);
}
Ok(out)
}
fn loo_aggregate(
values: &[Vec<f64>],
kind: PanelAggregator,
n: usize,
t: usize,
) -> Result<Vec<Vec<f64>>> {
if matches!(kind, PanelAggregator::Std) && n < 3 {
return Err(ForecastError::InsufficientData {
needed: 3,
got: n,
hint: Some("Std with exclude_self needs ≥ 3 series".into()),
});
}
let mut out = vec![vec![0.0_f64; t]; n];
let mut buf = vec![0.0_f64; n - 1];
for j in 0..t {
for i in 0..n {
let mut k = 0;
for (q, series) in values.iter().enumerate() {
if q != i {
buf[k] = series[j];
k += 1;
}
}
out[i][j] = compute_kind(&buf, kind);
}
}
Ok(out)
}
fn rank_per_series(values: &[Vec<f64>], t: usize) -> Vec<Vec<f64>> {
let n = values.len();
let mut out = vec![vec![0.0_f64; t]; n];
if n <= 1 {
if n == 1 {
out[0].iter_mut().for_each(|x| *x = 0.5);
}
return out;
}
let denom = (n - 1) as f64;
for j in 0..t {
for i in 0..n {
let xi = values[i][j];
let mut less = 0usize;
let mut equal = 0usize;
for series in values.iter() {
let xj = series[j];
if xj < xi {
less += 1;
} else if xj == xi {
equal += 1;
}
}
out[i][j] = ((less as f64) + 0.5 * (equal as f64 - 1.0)) / denom;
}
}
out
}
fn compute_kind(xs: &[f64], kind: PanelAggregator) -> f64 {
let n = xs.len();
if n == 0 {
return 0.0;
}
match kind {
PanelAggregator::Mean => xs.iter().sum::<f64>() / n as f64,
PanelAggregator::Median => {
let mut v = xs.to_vec();
v.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
if n % 2 == 1 {
v[n / 2]
} else {
0.5 * (v[n / 2 - 1] + v[n / 2])
}
}
PanelAggregator::Std => {
if n < 2 {
0.0
} else {
let m = xs.iter().sum::<f64>() / n as f64;
let var: f64 = xs.iter().map(|x| (x - m) * (x - m)).sum::<f64>() / (n - 1) as f64;
var.sqrt()
}
}
PanelAggregator::Rank => unreachable!("Rank handled by rank_per_series"),
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
fn make_panel() -> Vec<Vec<f64>> {
vec![
vec![1.0, 2.0, 3.0, 4.0],
vec![10.0, 20.0, 30.0, 40.0],
vec![100.0, 200.0, 300.0, 400.0],
]
}
#[test]
fn mean_full_panel_is_identical_per_series() {
let panel = make_panel();
let out = panel_mean(&panel, false).unwrap();
let expected = vec![37.0, 74.0, 111.0, 148.0];
for row in &out {
assert_eq!(row, &expected);
}
}
#[test]
fn mean_exclude_self_loo() {
let panel = make_panel();
let out = panel_mean(&panel, true).unwrap();
assert_relative_eq!(out[0][0], 55.0, epsilon = 1e-12);
assert_relative_eq!(out[1][0], 50.5, epsilon = 1e-12);
assert_relative_eq!(out[2][0], 5.5, epsilon = 1e-12);
}
#[test]
fn median_full_panel() {
let panel = make_panel();
let out = panel_median(&panel, false).unwrap();
assert_relative_eq!(out[0][0], 10.0, epsilon = 1e-12);
assert_relative_eq!(out[0][1], 20.0, epsilon = 1e-12);
}
#[test]
fn std_requires_at_least_two_series() {
let single = vec![vec![1.0, 2.0]];
let err = panel_std(&single, false).unwrap_err();
assert!(matches!(err, ForecastError::InsufficientData { .. }));
}
#[test]
fn std_full_panel() {
let panel = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let out = panel_std(&panel, false).unwrap();
assert_relative_eq!(out[0][0], 2.0_f64.sqrt(), epsilon = 1e-10);
}
#[test]
fn rank_recovers_ordering() {
let panel = make_panel();
let ranks = panel_rank(&panel).unwrap();
assert_relative_eq!(ranks[0][0], 0.0, epsilon = 1e-12);
assert_relative_eq!(ranks[1][0], 0.5, epsilon = 1e-12);
assert_relative_eq!(ranks[2][0], 1.0, epsilon = 1e-12);
}
#[test]
fn rank_handles_ties() {
let panel = vec![vec![5.0], vec![5.0], vec![10.0]];
let ranks = panel_rank(&panel).unwrap();
assert_relative_eq!(ranks[0][0], 0.25, epsilon = 1e-12);
assert_relative_eq!(ranks[1][0], 0.25, epsilon = 1e-12);
assert_relative_eq!(ranks[2][0], 1.0, epsilon = 1e-12);
}
#[test]
fn rejects_mismatched_lengths() {
let panel = vec![vec![1.0, 2.0], vec![3.0]];
let err = panel_mean(&panel, false).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
#[test]
fn rejects_empty_panel() {
let panel: Vec<Vec<f64>> = Vec::new();
let err = panel_mean(&panel, false).unwrap_err();
assert!(matches!(err, ForecastError::EmptyData));
}
#[test]
fn loo_requires_at_least_two_series() {
let single = vec![vec![1.0, 2.0]];
let err = panel_mean(&single, true).unwrap_err();
assert!(matches!(err, ForecastError::InsufficientData { .. }));
}
#[test]
fn loo_std_requires_at_least_three_series() {
let two = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let err = panel_std(&two, true).unwrap_err();
assert!(matches!(err, ForecastError::InsufficientData { .. }));
}
}