Skip to main content

solow_stats/
wald.rs

1//! Generic Wald and F tests for linear restrictions on a parameter vector.
2//!
3//! Both take a fitted parameter vector `θ̂`, its estimated covariance
4//! `V̂`, and a linear restriction `R θ = r`. `R` is `(q × k)` (one row
5//! per restriction), `r` is a `q`-vector.
6//!
7//! * [`wald_test`] — asymptotic Wald statistic `(R θ̂ - r)' (R V̂ R')⁻¹
8//!   (R θ̂ - r)`, distributed χ²(q) under the null.
9//! * [`f_test`] — same numerator, divided by `q`, referenced against the
10//!   `F(q, df_denom)` distribution. This is the small-sample-corrected
11//!   version usually reported for OLS.
12//!
13//! Both use a numerically-stable Cholesky-based inversion on `R V̂ R'`
14//! and return the statistic together with its degrees of freedom and
15//! p-value.
16
17use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
18use solow_core::{Error, Result};
19
20/// Result of a linear-restriction Wald test.
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22#[derive(Clone, Debug, PartialEq)]
23pub struct WaldResult {
24    /// Wald statistic `(R θ̂ - r)' (R V̂ R')⁻¹ (R θ̂ - r)`.
25    pub statistic: f64,
26    /// Upper-tail χ² p-value under `df` degrees of freedom.
27    pub p_value: f64,
28    /// Number of restrictions (rows of `R`).
29    pub df: usize,
30    /// `R θ̂ - r` — the fitted deviation from the restriction.
31    pub restriction_gap: Array1<f64>,
32}
33
34/// Result of a linear-restriction F test.
35#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
36#[derive(Clone, Debug, PartialEq)]
37pub struct FTestResult {
38    /// F statistic `(R θ̂ - r)' (R V̂ R')⁻¹ (R θ̂ - r) / q`.
39    pub statistic: f64,
40    /// Upper-tail F(df_num, df_denom) p-value.
41    pub p_value: f64,
42    /// Numerator degrees of freedom (number of restrictions).
43    pub df_num: usize,
44    /// Denominator degrees of freedom.
45    pub df_denom: f64,
46    /// `R θ̂ - r` — the fitted deviation from the restriction.
47    pub restriction_gap: Array1<f64>,
48}
49
50/// Wald test on a linear restriction `R θ = r`.
51///
52/// * `theta` — the fitted parameter vector `θ̂`.
53/// * `cov` — its estimated covariance matrix `V̂`.
54/// * `r_matrix` — the `(q × k)` restriction matrix.
55/// * `r_value` — the `q`-vector `r`.
56pub fn wald_test(
57    theta: ArrayView1<'_, f64>,
58    cov: ArrayView2<'_, f64>,
59    r_matrix: ArrayView2<'_, f64>,
60    r_value: ArrayView1<'_, f64>,
61) -> Result<WaldResult> {
62    let k = theta.len();
63    let (q, r_cols) = (r_matrix.nrows(), r_matrix.ncols());
64    if cov.nrows() != k || cov.ncols() != k {
65        return Err(Error::Shape(format!(
66            "wald_test: theta has {k} entries but cov is {}x{}",
67            cov.nrows(),
68            cov.ncols()
69        )));
70    }
71    if r_cols != k {
72        return Err(Error::Shape(format!(
73            "wald_test: r_matrix has {r_cols} columns but theta has {k} entries"
74        )));
75    }
76    if q == 0 || r_value.len() != q {
77        return Err(Error::Shape(format!(
78            "wald_test: r_matrix has {q} rows but r_value has {} entries",
79            r_value.len()
80        )));
81    }
82    let gap = compute_gap(theta, r_matrix, r_value);
83    let stat = quadratic_form(&gap, r_matrix, cov)?;
84    let p = chi2_upper_tail(stat, q as f64);
85    Ok(WaldResult {
86        statistic: stat,
87        p_value: p,
88        df: q,
89        restriction_gap: gap,
90    })
91}
92
93/// F test on the same restriction `R θ = r`, with `df_denom` denominator
94/// degrees of freedom (typically `n - k` for OLS).
95pub fn f_test(
96    theta: ArrayView1<'_, f64>,
97    cov: ArrayView2<'_, f64>,
98    r_matrix: ArrayView2<'_, f64>,
99    r_value: ArrayView1<'_, f64>,
100    df_denom: f64,
101) -> Result<FTestResult> {
102    if !(df_denom > 0.0 && df_denom.is_finite()) {
103        return Err(Error::Value(format!(
104            "f_test: df_denom must be finite and > 0 (got {df_denom})"
105        )));
106    }
107    let wald = wald_test(theta, cov, r_matrix, r_value)?;
108    let q = wald.df as f64;
109    let f_stat = wald.statistic / q;
110    Ok(FTestResult {
111        statistic: f_stat,
112        p_value: f_upper_tail(f_stat, q, df_denom),
113        df_num: wald.df,
114        df_denom,
115        restriction_gap: wald.restriction_gap,
116    })
117}
118
119fn compute_gap(
120    theta: ArrayView1<'_, f64>,
121    r_matrix: ArrayView2<'_, f64>,
122    r_value: ArrayView1<'_, f64>,
123) -> Array1<f64> {
124    let q = r_matrix.nrows();
125    let mut gap = Array1::<f64>::zeros(q);
126    for i in 0..q {
127        let mut s = 0.0_f64;
128        for j in 0..theta.len() {
129            s += r_matrix[[i, j]] * theta[j];
130        }
131        gap[i] = s - r_value[i];
132    }
133    gap
134}
135
136fn quadratic_form(
137    gap: &Array1<f64>,
138    r_matrix: ArrayView2<'_, f64>,
139    cov: ArrayView2<'_, f64>,
140) -> Result<f64> {
141    let q = gap.len();
142    let k = cov.nrows();
143    // R V R' — an (q x q) symmetric matrix. Build by RV then multiply again by R'.
144    let mut rv: Array2<f64> = Array2::zeros((q, k));
145    for i in 0..q {
146        for j in 0..k {
147            let mut s = 0.0_f64;
148            for a in 0..k {
149                s += r_matrix[[i, a]] * cov[[a, j]];
150            }
151            rv[[i, j]] = s;
152        }
153    }
154    let mut rvr: Vec<Vec<f64>> = vec![vec![0.0; q]; q];
155    for i in 0..q {
156        for j in 0..q {
157            let mut s = 0.0_f64;
158            for a in 0..k {
159                s += rv[[i, a]] * r_matrix[[j, a]];
160            }
161            rvr[i][j] = s;
162        }
163    }
164    // Solve (R V R') x = gap by Cholesky-with-fallback Gauss-Jordan.
165    let x = solve_symmetric(&mut rvr, gap.as_slice().unwrap())?;
166    let mut stat = 0.0_f64;
167    for i in 0..q {
168        stat += gap[i] * x[i];
169    }
170    Ok(stat)
171}
172
173fn solve_symmetric(m: &mut [Vec<f64>], rhs: &[f64]) -> Result<Vec<f64>> {
174    let n = m.len();
175    let mut a: Vec<Vec<f64>> = (0..n)
176        .map(|i| {
177            let mut row = Vec::with_capacity(n + 1);
178            row.extend_from_slice(&m[i]);
179            row.push(rhs[i]);
180            row
181        })
182        .collect();
183    for i in 0..n {
184        let mut pivot = i;
185        let mut best = a[i][i].abs();
186        for r in (i + 1)..n {
187            if a[r][i].abs() > best {
188                best = a[r][i].abs();
189                pivot = r;
190            }
191        }
192        if best < 1e-300 {
193            return Err(Error::Value(
194                "wald_test: R V R' is singular; cannot solve the restriction system".into(),
195            ));
196        }
197        if pivot != i {
198            a.swap(i, pivot);
199        }
200        let piv = a[i][i];
201        for c in 0..(n + 1) {
202            a[i][c] /= piv;
203        }
204        for r in 0..n {
205            if r == i {
206                continue;
207            }
208            let factor = a[r][i];
209            if factor == 0.0 {
210                continue;
211            }
212            for c in 0..(n + 1) {
213                a[r][c] -= factor * a[i][c];
214            }
215        }
216    }
217    Ok((0..n).map(|i| a[i][n]).collect())
218}
219
220// ---------------------------------------------------------------------------
221// χ² and F upper-tail p-values (local — kept private to keep the crate free
222// of a heavy stats-distributions dep at this leaf).
223// ---------------------------------------------------------------------------
224
225fn chi2_upper_tail(x: f64, df: f64) -> f64 {
226    if x <= 0.0 {
227        return 1.0;
228    }
229    if !x.is_finite() {
230        return 0.0;
231    }
232    reg_upper_gamma(0.5 * df, 0.5 * x).clamp(0.0, 1.0)
233}
234
235fn f_upper_tail(f: f64, df1: f64, df2: f64) -> f64 {
236    if f <= 0.0 {
237        return 1.0;
238    }
239    if !f.is_finite() {
240        return 0.0;
241    }
242    let x = df2 / (df2 + df1 * f);
243    reg_incomplete_beta(x, 0.5 * df2, 0.5 * df1).clamp(0.0, 1.0)
244}
245
246fn reg_upper_gamma(a: f64, x: f64) -> f64 {
247    if x < a + 1.0 {
248        1.0 - lower_series(a, x)
249    } else {
250        upper_cf(a, x)
251    }
252}
253
254fn lower_series(a: f64, x: f64) -> f64 {
255    const MAXIT: usize = 512;
256    const EPS: f64 = 3e-16;
257    let mut ap = a;
258    let mut sum = 1.0 / a;
259    let mut del = sum;
260    for _ in 0..MAXIT {
261        ap += 1.0;
262        del *= x / ap;
263        sum += del;
264        if del.abs() < sum.abs() * EPS {
265            break;
266        }
267    }
268    sum * (-x + a * x.ln() - ln_gamma(a)).exp()
269}
270
271fn upper_cf(a: f64, x: f64) -> f64 {
272    const MAXIT: usize = 512;
273    const EPS: f64 = 3e-16;
274    const FPMIN: f64 = 1e-300;
275    let mut b = x + 1.0 - a;
276    let mut c = 1.0 / FPMIN;
277    let mut d = 1.0 / b;
278    let mut h = d;
279    for i in 1..=MAXIT {
280        let an = -(i as f64) * (i as f64 - a);
281        b += 2.0;
282        d = an * d + b;
283        if d.abs() < FPMIN {
284            d = FPMIN;
285        }
286        c = b + an / c;
287        if c.abs() < FPMIN {
288            c = FPMIN;
289        }
290        d = 1.0 / d;
291        let del = d * c;
292        h *= del;
293        if (del - 1.0).abs() < EPS {
294            break;
295        }
296    }
297    h * (-x + a * x.ln() - ln_gamma(a)).exp()
298}
299
300fn reg_incomplete_beta(x: f64, a: f64, b: f64) -> f64 {
301    if x <= 0.0 {
302        return 0.0;
303    }
304    if x >= 1.0 {
305        return 1.0;
306    }
307    let bt = (ln_gamma(a + b) - ln_gamma(a) - ln_gamma(b) + a * x.ln() + b * (1.0 - x).ln()).exp();
308    if x < (a + 1.0) / (a + b + 2.0) {
309        bt * betacf(x, a, b) / a
310    } else {
311        1.0 - bt * betacf(1.0 - x, b, a) / b
312    }
313}
314
315fn betacf(x: f64, a: f64, b: f64) -> f64 {
316    const MAXIT: usize = 512;
317    const EPS: f64 = 3e-16;
318    const FPMIN: f64 = 1e-300;
319    let qab = a + b;
320    let qap = a + 1.0;
321    let qam = a - 1.0;
322    let mut c = 1.0;
323    let mut d = 1.0 - qab * x / qap;
324    if d.abs() < FPMIN {
325        d = FPMIN;
326    }
327    d = 1.0 / d;
328    let mut h = d;
329    for m in 1..=MAXIT {
330        let m_f = m as f64;
331        let m2 = 2.0 * m_f;
332        let aa = m_f * (b - m_f) * x / ((qam + m2) * (a + m2));
333        d = 1.0 + aa * d;
334        if d.abs() < FPMIN {
335            d = FPMIN;
336        }
337        c = 1.0 + aa / c;
338        if c.abs() < FPMIN {
339            c = FPMIN;
340        }
341        d = 1.0 / d;
342        h *= d * c;
343        let aa = -(a + m_f) * (qab + m_f) * x / ((a + m2) * (qap + m2));
344        d = 1.0 + aa * d;
345        if d.abs() < FPMIN {
346            d = FPMIN;
347        }
348        c = 1.0 + aa / c;
349        if c.abs() < FPMIN {
350            c = FPMIN;
351        }
352        d = 1.0 / d;
353        let del = d * c;
354        h *= del;
355        if (del - 1.0).abs() < EPS {
356            break;
357        }
358    }
359    h
360}
361
362fn ln_gamma(x: f64) -> f64 {
363    const G: f64 = 7.0;
364    const COEFF: [f64; 9] = [
365        0.999_999_999_999_809_93,
366        676.520_368_121_885_1,
367        -1_259.139_216_722_402_8,
368        771.323_428_777_653_1,
369        -176.615_029_162_140_59,
370        12.507_343_278_686_905,
371        -0.138_571_095_265_720_12,
372        9.984_369_578_019_571e-6,
373        1.505_632_735_149_311_6e-7,
374    ];
375    if x < 0.5 {
376        std::f64::consts::PI.ln() - (std::f64::consts::PI * x).sin().ln() - ln_gamma(1.0 - x)
377    } else {
378        let x = x - 1.0;
379        let mut a = COEFF[0];
380        for (i, &c) in COEFF.iter().enumerate().skip(1) {
381            a += c / (x + i as f64);
382        }
383        let t = x + G + 0.5;
384        0.5 * (2.0 * std::f64::consts::PI).ln() + (x + 0.5) * t.ln() - t + a.ln()
385    }
386}
387
388#[cfg(test)]
389mod tests {
390    use super::*;
391    use ndarray::array;
392
393    #[test]
394    fn wald_zero_when_restriction_is_satisfied() {
395        let theta = array![1.0, 2.0, 3.0];
396        let cov = Array2::<f64>::eye(3);
397        // R = [1, -1, 0], r = -1 → gap = 1 - 2 - (-1) = 0.
398        let r_matrix = array![[1.0, -1.0, 0.0]];
399        let r_value = array![-1.0];
400        let w = wald_test(theta.view(), cov.view(), r_matrix.view(), r_value.view()).unwrap();
401        assert!(w.statistic.abs() < 1e-12);
402        assert!((w.p_value - 1.0).abs() < 1e-12);
403        assert_eq!(w.df, 1);
404    }
405
406    #[test]
407    fn wald_large_when_restriction_is_far() {
408        let theta = array![5.0, 0.0];
409        let cov = Array2::<f64>::eye(2);
410        let r_matrix = array![[1.0, 0.0]];
411        let r_value = array![0.0];
412        let w = wald_test(theta.view(), cov.view(), r_matrix.view(), r_value.view()).unwrap();
413        // stat = 25 / 1 = 25 → tiny p.
414        assert!((w.statistic - 25.0).abs() < 1e-12);
415        assert!(w.p_value < 1e-6);
416    }
417
418    #[test]
419    fn f_test_reduces_to_wald_over_q_for_one_restriction() {
420        let theta = array![1.5, 0.0];
421        let cov = Array2::<f64>::eye(2);
422        let r_matrix = array![[1.0, 0.0]];
423        let r_value = array![0.0];
424        let w = wald_test(theta.view(), cov.view(), r_matrix.view(), r_value.view()).unwrap();
425        let f = f_test(
426            theta.view(),
427            cov.view(),
428            r_matrix.view(),
429            r_value.view(),
430            100.0,
431        )
432        .unwrap();
433        assert!((f.statistic - w.statistic).abs() < 1e-12);
434        assert_eq!(f.df_num, 1);
435        assert!((f.df_denom - 100.0).abs() < 1e-12);
436    }
437}