1use super::{BayesianConfig, BayesianFosrResult};
38use crate::error::FdarError;
39use crate::linalg::{cholesky_factor, cholesky_forward_back, compute_xtx};
40use crate::matrix::FdMatrix;
41use crate::regression::fdata_to_pc_1d;
42use rand::rngs::StdRng;
43use rand::{Rng, SeedableRng};
44use rand_distr::{Distribution, Gamma, StandardNormal};
45
46fn back_solve_lt(l: &[f64], z: &[f64], p: usize) -> Vec<f64> {
49 let mut v = z.to_vec();
50 for j in (0..p).rev() {
51 for k in (j + 1)..p {
52 v[j] -= l[k * p + j] * v[k];
54 }
55 v[j] /= l[j * p + j];
56 }
57 v
58}
59
60fn quantile_sorted(sorted: &[f64], q: f64) -> f64 {
62 let n = sorted.len();
63 if n == 0 {
64 return f64::NAN;
65 }
66 if n == 1 {
67 return sorted[0];
68 }
69 let pos = q * (n as f64 - 1.0);
70 let lo = pos.floor() as usize;
71 let hi = pos.ceil() as usize;
72 let frac = pos - lo as f64;
73 sorted[lo] * (1.0 - frac) + sorted[hi] * frac
74}
75
76#[must_use = "expensive computation whose result should not be discarded"]
101pub fn bayesian_fosr(
102 data: &FdMatrix,
103 predictors: &FdMatrix,
104 argvals: &[f64],
105 config: &BayesianConfig,
106) -> Result<BayesianFosrResult, FdarError> {
107 let (n, m_t) = data.shape();
108 let p = predictors.ncols();
109
110 if n < 2 || m_t == 0 || predictors.nrows() != n {
112 return Err(FdarError::InvalidDimension {
113 parameter: "data/predictors",
114 expected: format!("n >= 2, m_t > 0, predictors.nrows() == n (n={n})"),
115 actual: format!(
116 "n={n}, m_t={m_t}, predictors.nrows()={}",
117 predictors.nrows()
118 ),
119 });
120 }
121 if argvals.len() != m_t {
122 return Err(FdarError::InvalidDimension {
123 parameter: "argvals",
124 expected: format!("length == data.ncols() = {m_t}"),
125 actual: format!("length = {}", argvals.len()),
126 });
127 }
128 if p == 0 {
129 return Err(FdarError::InvalidDimension {
130 parameter: "predictors",
131 expected: "at least 1 predictor column".to_string(),
132 actual: "0 columns".to_string(),
133 });
134 }
135 if config.ncomp == 0 {
136 return Err(FdarError::InvalidParameter {
137 parameter: "ncomp",
138 message: "must be >= 1".to_string(),
139 });
140 }
141 if config.tau2 <= 0.0 {
142 return Err(FdarError::InvalidParameter {
143 parameter: "tau2",
144 message: format!("must be > 0, got {}", config.tau2),
145 });
146 }
147 if config.ig_a0 <= 0.0 || config.ig_b0 <= 0.0 {
148 return Err(FdarError::InvalidParameter {
149 parameter: "ig_a0/ig_b0",
150 message: format!("must be > 0, got a0={}, b0={}", config.ig_a0, config.ig_b0),
151 });
152 }
153 if config.n_iter == 0 {
154 return Err(FdarError::InvalidParameter {
155 parameter: "n_iter",
156 message: "must be >= 1".to_string(),
157 });
158 }
159 if config.thin == 0 {
160 return Err(FdarError::InvalidParameter {
161 parameter: "thin",
162 message: "must be >= 1".to_string(),
163 });
164 }
165
166 let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
168 let k = fpca.scores.ncols(); let mut xbar = vec![0.0f64; p];
175 for j in 0..p {
176 let col = predictors.column(j);
177 xbar[j] = col.iter().sum::<f64>() / n as f64;
178 }
179 let mut xc = FdMatrix::zeros(n, p);
181 for j in 0..p {
182 let src = predictors.column(j);
183 let dst = xc.column_mut(j);
184 for i in 0..n {
185 dst[i] = src[i] - xbar[j];
186 }
187 }
188
189 let xtx = compute_xtx(&xc); let mut xt_xi: Vec<Vec<f64>> = vec![vec![0.0f64; p]; k]; for kk in 0..k {
193 let score_col = fpca.scores.column(kk);
194 for j in 0..p {
195 let xj = xc.column(j);
196 let mut s = 0.0f64;
197 for i in 0..n {
198 s += xj[i] * score_col[i];
199 }
200 xt_xi[kk][j] = s;
201 }
202 }
203
204 let mut b_state: Vec<Vec<f64>> = vec![vec![0.0f64; p]; k];
206 let mut sigma2: Vec<f64> = vec![1.0f64; k];
207
208 let inv_tau2 = 1.0 / config.tau2;
209 let a_post_shape = config.ig_a0 + n as f64 / 2.0;
210
211 let mut rng = StdRng::seed_from_u64(config.seed);
212
213 let total = config.burn_in + config.n_iter * config.thin;
214 let q_retained = config.n_iter; let mut beta_draws: Vec<Vec<f64>> = vec![Vec::with_capacity(q_retained); p * m_t];
218 let mut beta_sum = vec![0.0f64; p * m_t];
220
221 for iter in 0..total {
222 for kk in 0..k {
223 let s2 = sigma2[kk];
224 let mut a = vec![0.0f64; p * p];
226 for idx in 0..p * p {
227 a[idx] = xtx[idx] / s2;
228 }
229 for d in 0..p {
230 a[d * p + d] += inv_tau2;
231 }
232 let l = cholesky_factor(&a, p)?;
233 let mut rhs = vec![0.0f64; p];
235 for j in 0..p {
236 rhs[j] = xt_xi[kk][j] / s2;
237 }
238 let mu_post = cholesky_forward_back(&l, &rhs, p);
239 let z: Vec<f64> = (0..p)
241 .map(|_| rng.sample::<f64, _>(StandardNormal))
242 .collect();
243 let v = back_solve_lt(&l, &z, p);
244 for j in 0..p {
245 b_state[kk][j] = mu_post[j] + v[j];
246 }
247 let score_col = fpca.scores.column(kk);
249 let mut rss = 0.0f64;
250 for i in 0..n {
251 let mut fit = 0.0f64;
252 for j in 0..p {
253 fit += xc.column(j)[i] * b_state[kk][j];
254 }
255 let r = score_col[i] - fit;
256 rss += r * r;
257 }
258 let rate = config.ig_b0 + rss / 2.0;
260 let gamma =
261 Gamma::new(a_post_shape, 1.0 / rate).map_err(|e| FdarError::ComputationFailed {
262 operation: "bayesian_fosr Inverse-Gamma draw",
263 detail: format!("Gamma::new failed (shape={a_post_shape}, rate={rate}): {e}"),
264 })?;
265 let g = gamma.sample(&mut rng);
266 sigma2[kk] = 1.0 / g.max(f64::MIN_POSITIVE);
267 }
268
269 if iter >= config.burn_in && (iter - config.burn_in) % config.thin == 0 {
271 for j in 0..p {
273 for t in 0..m_t {
274 let mut beta = 0.0f64;
275 for kk in 0..k {
276 beta += b_state[kk][j] * fpca.rotation[(t, kk)];
277 }
278 beta_draws[j * m_t + t].push(beta);
279 beta_sum[j * m_t + t] += beta;
280 }
281 }
282 }
283 }
284
285 let q = beta_draws[0].len().max(1);
286
287 let mut beta_mean = FdMatrix::zeros(p, m_t);
289 let mut beta_lower = FdMatrix::zeros(p, m_t);
290 let mut beta_upper = FdMatrix::zeros(p, m_t);
291 for j in 0..p {
292 for t in 0..m_t {
293 let cell = &mut beta_draws[j * m_t + t];
294 beta_mean[(j, t)] = beta_sum[j * m_t + t] / q as f64;
295 cell.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
296 beta_lower[(j, t)] = quantile_sorted(cell, 0.025);
297 beta_upper[(j, t)] = quantile_sorted(cell, 0.975);
298 }
299 }
300
301 let mut fitted = FdMatrix::zeros(n, m_t);
304 let mut residuals = FdMatrix::zeros(n, m_t);
305 let mut sigma2_mean = vec![0.0f64; m_t];
306 for t in 0..m_t {
307 for i in 0..n {
308 let mut val = fpca.mean[t];
309 for j in 0..p {
310 val += xc.column(j)[i] * beta_mean[(j, t)];
311 }
312 fitted[(i, t)] = val;
313 let r = data[(i, t)] - val;
314 residuals[(i, t)] = r;
315 sigma2_mean[t] += r * r;
316 }
317 sigma2_mean[t] /= n as f64;
318 }
319
320 Ok(BayesianFosrResult {
321 beta_mean,
322 beta_lower,
323 beta_upper,
324 fitted,
325 residuals,
326 sigma2_mean,
327 n_iter: config.n_iter,
328 burn_in: config.burn_in,
329 thin: config.thin,
330 ncomp: k,
331 })
332}
333
334#[cfg(test)]
335mod tests {
336 use super::*;
337 use crate::test_helpers::uniform_grid;
338 use std::f64::consts::PI;
339
340 fn default_config() -> BayesianConfig {
341 BayesianConfig {
342 ncomp: 4,
343 tau2: 100.0,
344 ig_a0: 0.001,
345 ig_b0: 0.001,
346 n_iter: 400,
347 burn_in: 200,
348 thin: 1,
349 seed: 20260824,
350 }
351 }
352
353 fn make_fosr_dataset(n: usize, m: usize) -> (FdMatrix, FdMatrix, Vec<f64>) {
355 let argvals = uniform_grid(m);
356 let x1: Vec<f64> = (0..n)
357 .map(|i| -1.0 + 2.0 * i as f64 / (n - 1).max(1) as f64)
358 .collect();
359 let predictors = FdMatrix::from_column_major(x1.clone(), n, 1).unwrap();
360
361 let mut y = vec![0.0f64; n * m];
362 for (t_idx, &tv) in argvals.iter().enumerate() {
363 let a_t = 0.5 * (2.0 * PI * tv).cos(); let beta_t = (PI * tv).sin(); for i in 0..n {
366 let noise = 0.02 * ((i as f64 * 1.2345 + t_idx as f64 * 0.678).sin());
367 y[i + t_idx * n] = a_t + x1[i] * beta_t + noise;
368 }
369 }
370 (
371 FdMatrix::from_column_major(y, n, m).unwrap(),
372 predictors,
373 argvals,
374 )
375 }
376
377 #[test]
378 fn bayesian_fosr_recovers_beta() {
379 let (data, predictors, argvals) = make_fosr_dataset(60, 25);
380 let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
381 assert_eq!(result.beta_mean.shape(), (1, 25));
382 let m = 25;
384 let mut dot = 0.0;
385 let mut nb = 0.0;
386 let mut nt = 0.0;
387 for t in 0..m {
388 let tv = argvals[t];
389 let truth = (PI * tv).sin();
390 let est = result.beta_mean[(0, t)];
391 dot += truth * est;
392 nb += est * est;
393 nt += truth * truth;
394 }
395 let corr = dot / (nb.sqrt() * nt.sqrt());
396 assert!(
397 corr > 0.9,
398 "posterior mean β should track the true coefficient (corr={corr:.3})"
399 );
400 }
401
402 #[test]
403 fn bayesian_fosr_credible_bands_bracket_mean() {
404 let (data, predictors, argvals) = make_fosr_dataset(50, 20);
405 let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
406 for t in 0..20 {
407 let lo = result.beta_lower[(0, t)];
408 let hi = result.beta_upper[(0, t)];
409 let mean = result.beta_mean[(0, t)];
410 assert!(lo <= mean + 1e-9, "lower band must be <= mean at t={t}");
411 assert!(hi >= mean - 1e-9, "upper band must be >= mean at t={t}");
412 assert!(lo.is_finite() && hi.is_finite());
413 }
414 }
415
416 #[test]
417 fn bayesian_fosr_is_deterministic_under_seed() {
418 let (data, predictors, argvals) = make_fosr_dataset(40, 15);
419 let cfg = default_config();
420 let r1 = bayesian_fosr(&data, &predictors, &argvals, &cfg).unwrap();
421 let r2 = bayesian_fosr(&data, &predictors, &argvals, &cfg).unwrap();
422 assert_eq!(
423 r1.beta_mean, r2.beta_mean,
424 "same seed → identical posterior mean"
425 );
426 assert_eq!(r1.beta_lower, r2.beta_lower);
427 assert_eq!(r1.beta_upper, r2.beta_upper);
428 assert_eq!(r1.sigma2_mean, r2.sigma2_mean);
429 }
430
431 #[test]
432 fn bayesian_fosr_sigma2_positive_and_shapes() {
433 let (data, predictors, argvals) = make_fosr_dataset(40, 18);
434 let result = bayesian_fosr(&data, &predictors, &argvals, &default_config()).unwrap();
435 assert_eq!(result.fitted.shape(), (40, 18));
436 assert_eq!(result.residuals.shape(), (40, 18));
437 assert_eq!(result.sigma2_mean.len(), 18);
438 assert!(result.sigma2_mean.iter().all(|&s| s > 0.0 && s.is_finite()));
439 }
440
441 #[test]
442 fn bayesian_fosr_errors_on_dimension_mismatch() {
443 let (data, _predictors, argvals) = make_fosr_dataset(30, 12);
444 let bad = FdMatrix::from_column_major(vec![0.0; 10], 10, 1).unwrap();
446 assert!(bayesian_fosr(&data, &bad, &argvals, &default_config()).is_err());
447 }
448
449 #[test]
450 fn bayesian_fosr_errors_on_invalid_params() {
451 let (data, predictors, argvals) = make_fosr_dataset(30, 12);
452 let mut cfg = default_config();
453 cfg.tau2 = -1.0;
454 assert!(bayesian_fosr(&data, &predictors, &argvals, &cfg).is_err());
455 let mut cfg2 = default_config();
456 cfg2.ncomp = 0;
457 assert!(bayesian_fosr(&data, &predictors, &argvals, &cfg2).is_err());
458 }
459}