1use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
18use solow_core::{Error, Result};
19
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22#[derive(Clone, Debug, PartialEq)]
23pub struct WaldResult {
24 pub statistic: f64,
26 pub p_value: f64,
28 pub df: usize,
30 pub restriction_gap: Array1<f64>,
32}
33
34#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
36#[derive(Clone, Debug, PartialEq)]
37pub struct FTestResult {
38 pub statistic: f64,
40 pub p_value: f64,
42 pub df_num: usize,
44 pub df_denom: f64,
46 pub restriction_gap: Array1<f64>,
48}
49
50pub 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
93pub 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 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 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
220fn 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 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 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}