use super::{BayesianConfig, BayesianFosrResult};
use crate::error::FdarError;
use crate::linalg::{cholesky_factor, cholesky_forward_back, compute_xtx};
use crate::matrix::FdMatrix;
use crate::regression::fdata_to_pc_1d;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
use rand_distr::{Distribution, Gamma, StandardNormal};
fn back_solve_lt(l: &[f64], z: &[f64], p: usize) -> Vec<f64> {
let mut v = z.to_vec();
for j in (0..p).rev() {
for k in (j + 1)..p {
v[j] -= l[k * p + j] * v[k];
}
v[j] /= l[j * p + j];
}
v
}
fn quantile_sorted(sorted: &[f64], q: f64) -> f64 {
let n = sorted.len();
if n == 0 {
return f64::NAN;
}
if n == 1 {
return sorted[0];
}
let pos = q * (n as f64 - 1.0);
let lo = pos.floor() as usize;
let hi = pos.ceil() as usize;
let frac = pos - lo as f64;
sorted[lo] * (1.0 - frac) + sorted[hi] * frac
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn bayesian_fosr(
data: &FdMatrix,
predictors: &FdMatrix,
argvals: &[f64],
config: &BayesianConfig,
) -> Result<BayesianFosrResult, FdarError> {
let (n, m_t) = data.shape();
let p = predictors.ncols();
if n < 2 || m_t == 0 || predictors.nrows() != n {
return Err(FdarError::InvalidDimension {
parameter: "data/predictors",
expected: format!("n >= 2, m_t > 0, predictors.nrows() == n (n={n})"),
actual: format!(
"n={n}, m_t={m_t}, predictors.nrows()={}",
predictors.nrows()
),
});
}
if argvals.len() != m_t {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: format!("length == data.ncols() = {m_t}"),
actual: format!("length = {}", argvals.len()),
});
}
if p == 0 {
return Err(FdarError::InvalidDimension {
parameter: "predictors",
expected: "at least 1 predictor column".to_string(),
actual: "0 columns".to_string(),
});
}
if config.ncomp == 0 {
return Err(FdarError::InvalidParameter {
parameter: "ncomp",
message: "must be >= 1".to_string(),
});
}
if config.tau2 <= 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "tau2",
message: format!("must be > 0, got {}", config.tau2),
});
}
if config.ig_a0 <= 0.0 || config.ig_b0 <= 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "ig_a0/ig_b0",
message: format!("must be > 0, got a0={}, b0={}", config.ig_a0, config.ig_b0),
});
}
if config.n_iter == 0 {
return Err(FdarError::InvalidParameter {
parameter: "n_iter",
message: "must be >= 1".to_string(),
});
}
if config.thin == 0 {
return Err(FdarError::InvalidParameter {
parameter: "thin",
message: "must be >= 1".to_string(),
});
}
let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
let k = fpca.scores.ncols();
let mut xbar = vec![0.0f64; p];
for j in 0..p {
let col = predictors.column(j);
xbar[j] = col.iter().sum::<f64>() / n as f64;
}
let mut xc = FdMatrix::zeros(n, p);
for j in 0..p {
let src = predictors.column(j);
let dst = xc.column_mut(j);
for i in 0..n {
dst[i] = src[i] - xbar[j];
}
}
let xtx = compute_xtx(&xc); let mut xt_xi: Vec<Vec<f64>> = vec![vec![0.0f64; p]; k]; for kk in 0..k {
let score_col = fpca.scores.column(kk);
for j in 0..p {
let xj = xc.column(j);
let mut s = 0.0f64;
for i in 0..n {
s += xj[i] * score_col[i];
}
xt_xi[kk][j] = s;
}
}
let mut b_state: Vec<Vec<f64>> = vec![vec![0.0f64; p]; k];
let mut sigma2: Vec<f64> = vec![1.0f64; k];
let inv_tau2 = 1.0 / config.tau2;
let a_post_shape = config.ig_a0 + n as f64 / 2.0;
let mut rng = StdRng::seed_from_u64(config.seed);
let total = config.burn_in + config.n_iter * config.thin;
let q_retained = config.n_iter;
let mut beta_draws: Vec<Vec<f64>> = vec![Vec::with_capacity(q_retained); p * m_t];
let mut beta_sum = vec![0.0f64; p * m_t];
for iter in 0..total {
for kk in 0..k {
let s2 = sigma2[kk];
let mut a = vec![0.0f64; p * p];
for idx in 0..p * p {
a[idx] = xtx[idx] / s2;
}
for d in 0..p {
a[d * p + d] += inv_tau2;
}
let l = cholesky_factor(&a, p)?;
let mut rhs = vec![0.0f64; p];
for j in 0..p {
rhs[j] = xt_xi[kk][j] / s2;
}
let mu_post = cholesky_forward_back(&l, &rhs, p);
let z: Vec<f64> = (0..p)
.map(|_| rng.sample::<f64, _>(StandardNormal))
.collect();
let v = back_solve_lt(&l, &z, p);
for j in 0..p {
b_state[kk][j] = mu_post[j] + v[j];
}
let score_col = fpca.scores.column(kk);
let mut rss = 0.0f64;
for i in 0..n {
let mut fit = 0.0f64;
for j in 0..p {
fit += xc.column(j)[i] * b_state[kk][j];
}
let r = score_col[i] - fit;
rss += r * r;
}
let rate = config.ig_b0 + rss / 2.0;
let gamma =
Gamma::new(a_post_shape, 1.0 / rate).map_err(|e| FdarError::ComputationFailed {
operation: "bayesian_fosr Inverse-Gamma draw",
detail: format!("Gamma::new failed (shape={a_post_shape}, rate={rate}): {e}"),
})?;
let g = gamma.sample(&mut rng);
sigma2[kk] = 1.0 / g.max(f64::MIN_POSITIVE);
}
if iter >= config.burn_in && (iter - config.burn_in) % config.thin == 0 {
for j in 0..p {
for t in 0..m_t {
let mut beta = 0.0f64;
for kk in 0..k {
beta += b_state[kk][j] * fpca.rotation[(t, kk)];
}
beta_draws[j * m_t + t].push(beta);
beta_sum[j * m_t + t] += beta;
}
}
}
}
let q = beta_draws[0].len().max(1);
let mut beta_mean = FdMatrix::zeros(p, m_t);
let mut beta_lower = FdMatrix::zeros(p, m_t);
let mut beta_upper = FdMatrix::zeros(p, m_t);
for j in 0..p {
for t in 0..m_t {
let cell = &mut beta_draws[j * m_t + t];
beta_mean[(j, t)] = beta_sum[j * m_t + t] / q as f64;
cell.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
beta_lower[(j, t)] = quantile_sorted(cell, 0.025);
beta_upper[(j, t)] = quantile_sorted(cell, 0.975);
}
}
let mut fitted = FdMatrix::zeros(n, m_t);
let mut residuals = FdMatrix::zeros(n, m_t);
let mut sigma2_mean = vec![0.0f64; m_t];
for t in 0..m_t {
for i in 0..n {
let mut val = fpca.mean[t];
for j in 0..p {
val += xc.column(j)[i] * beta_mean[(j, t)];
}
fitted[(i, t)] = val;
let r = data[(i, t)] - val;
residuals[(i, t)] = r;
sigma2_mean[t] += r * r;
}
sigma2_mean[t] /= n as f64;
}
Ok(BayesianFosrResult {
beta_mean,
beta_lower,
beta_upper,
fitted,
residuals,
sigma2_mean,
n_iter: config.n_iter,
burn_in: config.burn_in,
thin: config.thin,
ncomp: k,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_helpers::uniform_grid;
use std::f64::consts::PI;
fn default_config() -> BayesianConfig {
BayesianConfig {
ncomp: 4,
tau2: 100.0,
ig_a0: 0.001,
ig_b0: 0.001,
n_iter: 400,
burn_in: 200,
thin: 1,
seed: 20260824,
}
}
fn make_fosr_dataset(n: usize, m: usize) -> (FdMatrix, FdMatrix, Vec<f64>) {
let argvals = uniform_grid(m);
let x1: Vec<f64> = (0..n)
.map(|i| -1.0 + 2.0 * i as f64 / (n - 1).max(1) as f64)
.collect();
let predictors = FdMatrix::from_column_major(x1.clone(), n, 1).unwrap();
let mut y = vec![0.0f64; n * m];
for (t_idx, &tv) in argvals.iter().enumerate() {
let a_t = 0.5 * (2.0 * PI * tv).cos(); let beta_t = (PI * tv).sin(); for i in 0..n {
let noise = 0.02 * ((i as f64 * 1.2345 + t_idx as f64 * 0.678).sin());
y[i + t_idx * n] = a_t + x1[i] * beta_t + noise;
}
}
(
FdMatrix::from_column_major(y, n, m).unwrap(),
predictors,
argvals,
)
}
#[test]
fn bayesian_fosr_recovers_beta() {
let (data, predictors, argvals) = make_fosr_dataset(60, 25);
let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
assert_eq!(result.beta_mean.shape(), (1, 25));
let m = 25;
let mut dot = 0.0;
let mut nb = 0.0;
let mut nt = 0.0;
for t in 0..m {
let tv = argvals[t];
let truth = (PI * tv).sin();
let est = result.beta_mean[(0, t)];
dot += truth * est;
nb += est * est;
nt += truth * truth;
}
let corr = dot / (nb.sqrt() * nt.sqrt());
assert!(
corr > 0.9,
"posterior mean β should track the true coefficient (corr={corr:.3})"
);
}
#[test]
fn bayesian_fosr_credible_bands_bracket_mean() {
let (data, predictors, argvals) = make_fosr_dataset(50, 20);
let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
for t in 0..20 {
let lo = result.beta_lower[(0, t)];
let hi = result.beta_upper[(0, t)];
let mean = result.beta_mean[(0, t)];
assert!(lo <= mean + 1e-9, "lower band must be <= mean at t={t}");
assert!(hi >= mean - 1e-9, "upper band must be >= mean at t={t}");
assert!(lo.is_finite() && hi.is_finite());
}
}
#[test]
fn bayesian_fosr_is_deterministic_under_seed() {
let (data, predictors, argvals) = make_fosr_dataset(40, 15);
let cfg = default_config();
let r1 = bayesian_fosr(&data, &predictors, &argvals, &cfg).unwrap();
let r2 = bayesian_fosr(&data, &predictors, &argvals, &cfg).unwrap();
assert_eq!(
r1.beta_mean, r2.beta_mean,
"same seed → identical posterior mean"
);
assert_eq!(r1.beta_lower, r2.beta_lower);
assert_eq!(r1.beta_upper, r2.beta_upper);
assert_eq!(r1.sigma2_mean, r2.sigma2_mean);
}
#[test]
fn bayesian_fosr_sigma2_positive_and_shapes() {
let (data, predictors, argvals) = make_fosr_dataset(40, 18);
let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
assert_eq!(result.fitted.shape(), (40, 18));
assert_eq!(result.residuals.shape(), (40, 18));
assert_eq!(result.sigma2_mean.len(), 18);
assert!(result.sigma2_mean.iter().all(|&s| s > 0.0 && s.is_finite()));
}
#[test]
fn bayesian_fosr_errors_on_dimension_mismatch() {
let (data, _predictors, argvals) = make_fosr_dataset(30, 12);
let bad = FdMatrix::from_column_major(vec![0.0; 10], 10, 1).unwrap();
assert!(bayesian_fosr(&data, &bad, &argvals, &default_config()).is_err());
}
#[test]
fn bayesian_fosr_errors_on_invalid_params() {
let (data, predictors, argvals) = make_fosr_dataset(30, 12);
let mut cfg = default_config();
cfg.tau2 = -1.0;
assert!(bayesian_fosr(&data, &predictors, &argvals, &cfg).is_err());
let mut cfg2 = default_config();
cfg2.ncomp = 0;
assert!(bayesian_fosr(&data, &predictors, &argvals, &cfg2).is_err());
}
}