Skip to main content

fdars_core/boosting_regression/
bayesian.rs

1//! Bayesian function-on-scalar regression via conjugate Gibbs sampler (REG-06-04).
2//!
3//! Fits `Y_i(t) = μ(t) + Σ_j x̃_ij · β_j(t) + ε_i(t)` where `x̃` are the
4//! mean-centered scalar predictors and `β_j(t)` are functional coefficients. The
5//! functional dimension is compressed with FPCA (`fdata_to_pc_1d`): the response
6//! is projected onto its top-`K` functional principal components, and for each
7//! component `k` the FPC scores are regressed on the predictors with a conjugate
8//! Normal / Inverse-Gamma Gibbs sampler. Each retained draw reconstructs the
9//! coefficient functions `β_j(t) = Σ_k b_{jk} · φ_k(t)` from the FPCA rotation, so
10//! posterior summaries (mean + **pointwise** credible bands) are obtained directly
11//! in the functional domain.
12//!
13//! # Full conditionals (per FPC component `k`)
14//!
15//! Model: `ξ_k = X̃ · b_k + ε_k`, `ξ_k ∈ Rⁿ` (scores of component k), `X̃ ∈ R^{n×p}`.
16//! Priors: `b_k | σ²_k ~ N(0, τ²·I_p)`, `σ²_k ~ IG(a₀, b₀)`.
17//!
18//! - `b_k | · ~ N(μ_post, A⁻¹)` with precision `A = X̃'X̃/σ²_k + I_p/τ²`,
19//!   `μ_post = A⁻¹ · X̃'ξ_k / σ²_k`. Sampled via Cholesky `A = LLᵀ`:
20//!   `b_k = μ_post + Lᵀ⁻¹ z`, `z ~ N(0, I_p)` (Rue 2001) — `Cov = (LLᵀ)⁻¹ = A⁻¹`.
21//! - `σ²_k | · ~ IG(a₀ + n/2, b₀ + RSS_k/2)`, drawn as `1 / Gamma(α, 1/β)`.
22//!
23//! The chain is fully deterministic given `config.seed`
24//! (`StdRng::seed_from_u64(seed)`).
25//!
26//! # References
27//!
28//! Rue (2001). Fast sampling of Gaussian Markov random fields. *JRSS-B* 63(2).
29//! Goldsmith et al. (2015). *JCGS* 23(1). Jiang et al. (2025), arXiv:2505.05633.
30//!
31//! # Divergences from refund
32//!
33//! `refund`'s Bayesian FOSR uses spline basis priors with random effects; this
34//! implementation uses FPCA score compression (`fdata_to_pc_1d`) for simplicity and
35//! zero new dependencies. Pointwise credible bands only (no simultaneous bands).
36
37use 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
46/// Solve `Lᵀ v = z` (back substitution) for a lower-triangular `L` stored flat
47/// row-major (`L[i*p + j]`, `i ≥ j`). Used to draw `N(0, (LLᵀ)⁻¹)` as `Lᵀ⁻¹ z`.
48fn 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            // (Lᵀ)_{jk} = L_{kj} = l[k*p + j]
53            v[j] -= l[k * p + j] * v[k];
54        }
55        v[j] /= l[j * p + j];
56    }
57    v
58}
59
60/// Linear-interpolated quantile of an already-sorted slice.
61fn 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/// Bayesian function-on-scalar regression via conjugate Gibbs sampler.
77///
78/// Compresses the functional response with FPCA, runs a conjugate Normal /
79/// Inverse-Gamma Gibbs sampler on the FPC-score regression coefficients, and
80/// reconstructs posterior summaries of the coefficient functions `β_j(t)`.
81///
82/// # Arguments
83///
84/// * `data` — Functional response Y (n × m_t, column-major).
85/// * `predictors` — Scalar predictor matrix X (n × p). Centered internally.
86/// * `argvals` — Response grid evaluation points (length m_t).
87/// * `config` — [`BayesianConfig`] (FPC components, priors, iterations, seed).
88///
89/// # Returns
90///
91/// [`BayesianFosrResult`] with posterior-mean coefficient functions `beta_mean`
92/// (p × m_t), pointwise 2.5% / 97.5% credible bands, posterior-mean fitted values
93/// and residuals, and the posterior-mean residual variance `sigma2_mean(t)`.
94///
95/// # Errors
96///
97/// [`FdarError::InvalidDimension`] on shape mismatch; [`FdarError::InvalidParameter`]
98/// on out-of-range config; [`FdarError::ComputationFailed`] if FPCA or a Cholesky
99/// draw fails.
100#[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    // ---- Validation --------------------------------------------------------
111    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    // ---- FPCA score compression of the response ----------------------------
167    let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
168    let k = fpca.scores.ncols(); // actual components retained (≤ ncomp)
169                                 // rotation: m_t × k loadings φ_k(t); scores: n × k ; mean: m_t response mean.
170
171    // ---- Center predictors -------------------------------------------------
172    // β_j(t) is the slope of the centered predictor; the response intercept μ(t)
173    // is carried by the FPCA mean function.
174    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    // Centered design X̃ (n × p, column-major).
180    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    // Precompute X̃'X̃ (p × p, row-major) and X̃'ξ_k (p-vector per component).
190    let xtx = compute_xtx(&xc); // p × p row-major
191    let mut xt_xi: Vec<Vec<f64>> = vec![vec![0.0f64; p]; k]; // [k][j] = Σ_i x̃_ij ξ_ik
192    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    // Gibbs state: b[k] ∈ R^p (init 0), sigma2[k] (init 1.0).
205    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; // number of retained draws
215
216    // Per-(j,t) storage of reconstructed β draws for pointwise quantiles.
217    let mut beta_draws: Vec<Vec<f64>> = vec![Vec::with_capacity(q_retained); p * m_t];
218    // Accumulator for posterior means.
219    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            // Precision A = X̃'X̃ / σ²_k + I_p / τ²  (row-major p×p).
225            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            // μ_post = A⁻¹ (X̃'ξ_k / σ²_k)
234            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            // Draw z ~ N(0, I_p) and v = Lᵀ⁻¹ z, then b_k = μ_post + v.
240            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            // RSS_k = ||ξ_k - X̃ b_k||²
248            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            // σ²_k ~ IG(a_post, b0 + RSS/2)  drawn as 1 / Gamma(shape, scale = 1/rate)
259            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        // Retain thinned post-burn-in draws.
270        if iter >= config.burn_in && (iter - config.burn_in) % config.thin == 0 {
271            // Reconstruct β_j(t) = Σ_k b_{kj} · φ_k(t) and record.
272            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    // ---- Posterior summaries ----------------------------------------------
288    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    // ---- Fitted / residuals / σ²(t) ---------------------------------------
302    // fitted(i,t) = mean_Y(t) + Σ_j x̃_ij · β̄_j(t)
303    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    /// Y_i(t) = a(t) + x_i · β(t) + small noise, β(t) = sin(π t).
354    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(); // intercept function
364            let beta_t = (PI * tv).sin(); // true coefficient
365            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        // Posterior mean β̄(t) correlates strongly with the true sin(π t).
383        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        // predictors with wrong row count
445        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}