use symplex::linprog::{q, qi};
use symplex::num_traits::{One, Zero};
use symplex::prelude::*;
use symplex::stats::data::{self, Ddof, from_i64};
use symplex::stats::hypothesis::Alternative;
use symplex::stats::regression::{
Design, LogitOpts, hat_matrix, logit, ols, polyfit, r_squared_from_correlation,
simple_linear_regression, slope_from_correlation, vif, wls,
};
fn close(actual: f64, expected: f64, tol: f64) {
assert!(
(actual - expected).abs() < tol,
"got {actual}, expected {expected} (tol {tol})"
);
}
fn ev(e: &Ex) -> f64 {
e.eval_f64().unwrap()
}
fn col(x: &[i64]) -> Vec<Vec<Q>> {
x.iter().map(|&v| vec![qi(v)]).collect()
}
fn rows(r: &[&[i64]]) -> Vec<Vec<Q>> {
r.iter().map(|row| from_i64(row)).collect()
}
fn qs(entries: &[(i64, i64)]) -> Vec<Q> {
entries.iter().map(|&(n, d)| q(n, d)).collect()
}
fn d1() -> (Vec<Q>, Vec<Q>) {
(
from_i64(&[1, 2, 3, 4, 5, 6, 7]),
from_i64(&[2, 3, 5, 4, 6, 8, 9]),
)
}
fn d2() -> (Vec<Vec<Q>>, Vec<Q>) {
(
rows(&[
&[1, 5],
&[2, 3],
&[3, 8],
&[4, 1],
&[5, 7],
&[6, 2],
&[7, 9],
&[8, 4],
]),
from_i64(&[6, 5, 10, 4, 12, 7, 15, 9]),
)
}
fn d3() -> (Vec<Vec<Q>>, Vec<Q>) {
(col(&[1, 2, 3, 4, 5]), from_i64(&[2, 3, 7, 8, 11]))
}
fn d4_weights() -> Vec<Q> {
from_i64(&[1, 2, 3, 1, 2, 3, 1])
}
#[test]
fn ols_simple_coefficients_fitted_residuals_exact() {
let (x, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
assert_eq!(fit.coefficients, vec![q(5, 7), q(8, 7)]);
assert_eq!(
fit.fitted,
qs(&[(13, 7), (3, 1), (29, 7), (37, 7), (45, 7), (53, 7), (61, 7)])
);
assert_eq!(
fit.residuals,
qs(&[(1, 7), (0, 1), (6, 7), (-9, 7), (-3, 7), (3, 7), (2, 7)])
);
assert_eq!(fit.nobs(), 7);
assert_eq!(fit.n_params(), 2);
assert!(fit.has_constant());
assert_eq!(fit.response(), &y[..]);
assert_eq!(fit.design().shape(), (7, 2));
assert_eq!(fit.design().col(1), x);
}
#[test]
fn ols_simple_sums_of_squares_and_r_squared_exact() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
assert_eq!(fit.ssr, q(20, 7));
assert_eq!(fit.ess, q(256, 7));
assert_eq!(fit.tss, q(276, 7));
assert_eq!(fit.r_squared, q(64, 69));
assert_eq!(fit.adjusted_r_squared, q(21, 23));
assert_eq!(fit.mse_resid, q(4, 7));
assert_eq!((fit.df_model, fit.df_resid), (1, 5));
close(
data::to_f64(std::slice::from_ref(&fit.r_squared))[0],
0.927_536_231_884_058,
1e-15,
);
}
#[test]
fn ols_simple_cov_params_and_standard_errors() {
let ctx = Context::new();
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
assert_eq!(
fit.normalized_cov_params().to_rows(),
vec![qs(&[(5, 7), (-1, 7)]), qs(&[(-1, 7), (1, 28)])]
);
assert_eq!(
fit.cov_params.to_rows(),
vec![qs(&[(20, 49), (-4, 49)]), qs(&[(-4, 49), (1, 49)])]
);
let se = fit.standard_errors(&ctx);
close(ev(&se[0]), 0.638_876_564_999_939_8, 1e-12);
assert_eq!(se[1].as_rational(), Some(q(1, 7)));
close(
ev(&fit.residual_standard_error(&ctx)),
0.755_928_946_018_454_4,
1e-12,
);
}
#[test]
fn ols_simple_t_statistics_and_p_values() {
let ctx = Context::new();
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let t = fit.t_statistics(&ctx).unwrap();
close(ev(&t[0]), 1.118_033_988_749_899, 1e-12);
assert_eq!(t[0].powi(2).simplify().as_rational(), Some(q(5, 4)));
assert_eq!(t[1].as_rational(), Some(qi(8)));
let p = fit.p_values(&ctx).unwrap();
close(ev(&p[0]), 0.314_372_637_647_015_45, 1e-12);
close(ev(&p[1]), 0.000_492_906_660_572_44, 1e-12);
}
#[test]
fn ols_simple_coefficient_tests_are_two_sided_student_t() {
let ctx = Context::new();
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let tests = fit.coefficient_tests(&ctx).unwrap();
assert_eq!(tests.len(), 2);
for t in &tests {
assert_eq!(t.alternative, Alternative::TwoSided);
assert_eq!(t.df.as_ref().and_then(Ex::as_rational), Some(qi(5)));
}
close(tests[1].statistic_f64().unwrap(), 8.0, 1e-12);
close(
tests[1].p_value_f64().unwrap(),
0.000_492_906_660_572_444_5,
1e-12,
);
}
#[test]
fn ols_simple_overall_f_test() {
let ctx = Context::new();
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
assert_eq!(fit.f_statistic().unwrap(), qi(64));
let f = fit.f_test(&ctx).unwrap();
assert_eq!(f.statistic_exact(), Some(qi(64)));
assert_eq!(f.alternative, Alternative::Greater);
assert_eq!(f.df.as_ref().and_then(Ex::as_rational), Some(qi(5)));
close(f.p_value_f64().unwrap(), 0.000_492_906_660_572_443_7, 1e-12);
let p = fit.p_values(&ctx).unwrap();
close(f.p_value_f64().unwrap(), ev(&p[1]), 1e-14);
}
#[test]
fn ols_simple_anova_table() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let t = fit.anova_table().unwrap();
assert_eq!(t.ss_model, q(256, 7));
assert_eq!(t.df_model, 1);
assert_eq!(t.ms_model, q(256, 7));
assert_eq!(t.ss_resid, q(20, 7));
assert_eq!(t.df_resid, 5);
assert_eq!(t.ms_resid, q(4, 7));
assert_eq!(t.ss_total, q(276, 7));
assert_eq!(t.df_total, 6);
assert_eq!(t.f, qi(64));
assert_eq!(&t.ss_model + &t.ss_resid, t.ss_total);
}
#[test]
fn ols_simple_conf_int_matches_statsmodels() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, -0.927_998_778_916_851_8, 1e-9);
close(ci[0].upper, 2.356_570_207_488_285_7, 1e-9);
close(ci[1].lower, 0.775_631_166_337_669_6, 1e-9);
close(ci[1].upper, 1.510_083_119_376_616_9, 1e-9);
let ci = fit.conf_int(0.90).unwrap();
close(ci[0].lower, -0.573_081_468_778_001_3, 1e-9);
close(ci[0].upper, 2.001_652_897_349_434_7, 1e-9);
close(ci[1].lower, 0.854_993_089_523_854_2, 1e-9);
close(ci[1].upper, 1.430_721_196_190_432_3, 1e-9);
assert!(fit.conf_int(1.0).is_err());
assert!(fit.conf_int(0.0).is_err());
}
#[test]
fn ols_simple_predict_and_intervals() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
assert_eq!(fit.predict(&[qi(8)]).unwrap(), q(69, 7));
let mean_ci = fit
.confidence_interval_mean_response(&[qi(8)], 0.95)
.unwrap();
close(mean_ci.lower, 8.214_858_363_940_294, 1e-9);
close(mean_ci.upper, 11.499_427_350_345_432, 1e-9);
let obs_ci = fit.prediction_interval(&[qi(8)], 0.95).unwrap();
close(obs_ci.lower, 7.312_926_660_379_568, 1e-9);
close(obs_ci.upper, 12.401_359_053_906_159, 1e-9);
assert_eq!(fit.predict(&[q(5, 2)]).unwrap(), q(25, 7));
let mean_ci = fit
.confidence_interval_mean_response(&[q(5, 2)], 0.95)
.unwrap();
close(mean_ci.lower, 2.653_363_630_129_891, 1e-9);
close(mean_ci.upper, 4.489_493_512_727_258, 1e-9);
let obs_ci = fit.prediction_interval(&[q(5, 2)], 0.95).unwrap();
close(obs_ci.lower, 1.422_293_644_137_870_4, 1e-9);
close(obs_ci.upper, 5.720_563_498_719_279_5, 1e-9);
assert!(fit.predict(&[qi(1), qi(2)]).is_err());
assert!(fit.prediction_interval(&[qi(8)], 1.5).is_err());
}
#[test]
fn ols_simple_leverage_cooks_distance_durbin_watson_exact() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let lev = fit.leverage();
assert_eq!(
lev,
qs(&[(13, 28), (2, 7), (5, 28), (1, 7), (5, 28), (2, 7), (13, 28)])
);
assert_eq!(lev.iter().fold(Q::zero(), |a, b| a + b), qi(2)); assert_eq!(
fit.cooks_distance().unwrap(),
qs(&[
(13, 450),
(0, 1),
(90, 529),
(9, 32),
(45, 1058),
(9, 100),
(26, 225)
])
);
assert_eq!(fit.durbin_watson().unwrap(), q(67, 28));
close(
data::to_f64(&[q(67, 28)])[0],
2.392_857_142_857_143_2,
1e-15,
);
}
#[test]
fn ols_simple_hat_matrix_is_symmetric_idempotent_projection() {
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
let h = fit.hat_matrix().unwrap();
assert_eq!(h.shape(), (7, 7));
assert!(h.is_symmetric());
assert_eq!(h.matmul(&h).unwrap(), h);
assert_eq!(h.trace().unwrap(), qi(2));
assert_eq!(h.diagonal(), fit.leverage());
let ycol = QMatrix::new(y.iter().map(|v| vec![v.clone()]).collect()).unwrap();
assert_eq!(h.matmul(&ycol).unwrap().col(0), fit.fitted);
assert_eq!(hat_matrix(fit.design()).unwrap(), h);
let bad = QMatrix::from_i64(&[&[1, 2], &[2, 4], &[3, 6]]).unwrap();
assert!(matches!(
hat_matrix(&bad),
Err(SymplexError::InvalidArgument { .. })
));
}
#[test]
fn ols_simple_log_likelihood_aic_bic() {
let ctx = Context::new();
let (_, y) = d1();
let fit = ols(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), true).unwrap();
close(
ev(&fit.log_likelihood(&ctx).unwrap()),
-6.796_261_646_484_484,
1e-12,
);
close(ev(&fit.aic(&ctx).unwrap()), 17.592_523_292_968_97, 1e-12);
close(ev(&fit.bic(&ctx).unwrap()), 17.484_343_591_079_593, 1e-12);
}
#[test]
fn simple_linear_regression_matches_scipy_linregress() {
let ctx = Context::new();
let (x, y) = d1();
let fit = simple_linear_regression(&x, &y).unwrap();
assert_eq!(fit.coefficients, vec![q(5, 7), q(8, 7)]);
assert_eq!(fit.r_squared, q(64, 69));
close(
data::to_f64(std::slice::from_ref(&fit.r_squared))[0].sqrt(),
0.963_086_824_686_153_6,
1e-15,
);
let se = fit.standard_errors(&ctx);
close(ev(&se[1]), 0.142_857_142_857_142_85, 1e-15);
close(ev(&se[0]), 0.638_876_564_999_939_9, 1e-12);
close(
ev(&fit.p_values(&ctx).unwrap()[1]),
0.000_492_906_660_572_444_5,
1e-12,
);
assert!(simple_linear_regression(&x[..6], &y).is_err());
}
#[test]
fn correlation_bridges_recover_r_squared_and_slope() {
let ctx = Context::new();
let (x, y) = d1();
let r = data::pearson(&ctx, &x, &y).unwrap();
close(ev(&r), 0.963_086_824_686_153_7, 1e-15);
let r2 = r_squared_from_correlation(&r);
assert_eq!(r2.as_rational(), Some(q(64, 69)));
let sx = data::std(&ctx, &x, Ddof::Sample).unwrap();
let sy = data::std(&ctx, &y, Ddof::Sample).unwrap();
let slope = slope_from_correlation(&r, &sx, &sy).unwrap();
close(ev(&slope), 8.0 / 7.0, 1e-15);
assert_eq!(slope.as_rational(), Some(q(8, 7)));
let sx0 = data::std(&ctx, &x, Ddof::Population).unwrap();
let sy0 = data::std(&ctx, &y, Ddof::Population).unwrap();
close(
ev(&slope_from_correlation(&r, &sx0, &sy0).unwrap()),
8.0 / 7.0,
1e-15,
);
assert!(slope_from_correlation(&r, &ctx.zero(), &sy).is_err());
}
#[test]
fn polyfit_degree_one_returns_highest_power_first() {
let (x, y) = d1();
assert_eq!(polyfit(&x, &y, 1).unwrap(), vec![q(5, 7), q(8, 7)]); assert_eq!(polyfit(&x, &y, 0).unwrap(), vec![q(37, 7)]);
}
#[test]
fn ols_two_regressors_coefficients_and_residuals_exact() {
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
assert_eq!(
fit.coefficients,
qs(&[(731, 3908), (13559, 19540), (5201, 4885)])
);
assert_eq!(
fit.fitted,
qs(&[
(60617, 9770),
(18637, 3908),
(52691, 4885),
(15739, 3908),
(108539, 9770),
(126617, 19540),
(71451, 4885),
(195343, 19540)
])
);
assert_eq!(
fit.residuals,
qs(&[
(-1997, 9770),
(903, 3908),
(-3841, 4885),
(-107, 3908),
(8701, 9770),
(10163, 19540),
(1824, 4885),
(-19483, 19540)
])
);
for j in 0..3 {
let c = fit.design().col(j);
let s = c
.iter()
.zip(&fit.residuals)
.fold(Q::zero(), |a, (u, e)| a + u * e);
assert!(s.is_zero());
}
}
#[test]
fn ols_two_regressors_fit_statistics_exact() {
let ctx = Context::new();
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
assert_eq!(fit.ssr, q(56889, 19540));
assert_eq!(fit.tss, qi(98));
assert_eq!(fit.ess, q(1858031, 19540));
assert_eq!(fit.r_squared, q(37919, 39080));
assert_eq!(fit.adjusted_r_squared, q(187273, 195400));
assert_eq!(fit.mse_resid, q(56889, 97700));
assert_eq!((fit.df_model, fit.df_resid), (2, 5));
assert_eq!(fit.f_statistic().unwrap(), q(189595, 2322));
let f = fit.f_test(&ctx).unwrap();
close(f.statistic_f64().unwrap(), 81.651_593_453_919_09, 1e-12);
close(f.p_value_f64().unwrap(), 0.000_152_122_747_883_555_3, 1e-12);
let t = fit.anova_table().unwrap();
assert_eq!(t.ms_model, q(1858031, 39080));
assert_eq!(t.df_total, 7);
}
#[test]
fn ols_two_regressors_cov_params_exact_and_bse() {
let ctx = Context::new();
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
assert_eq!(
fit.cov_params.to_rows(),
vec![
qs(&[
(205198623, 381811600),
(-21674709, 381811600),
(-1024002, 23863225)
]),
qs(&[
(-21674709, 381811600),
(26794719, 1909058000),
(-625779, 477264500)
]),
qs(&[
(-1024002, 23863225),
(-625779, 477264500),
(1194669, 119316125)
]),
]
);
assert!(fit.cov_params.is_symmetric());
let se = fit.standard_errors(&ctx);
close(ev(&se[0]), 0.733_099_058_342_342_8, 1e-12);
close(ev(&se[1]), 0.118_471_814_987_821_83, 1e-12);
close(ev(&se[2]), 0.100_063_163_046_064_39, 1e-12);
close(
ev(&fit.residual_standard_error(&ctx)),
0.763_074_372_155_916_7,
1e-12,
);
}
#[test]
fn ols_two_regressors_t_and_p_values() {
let ctx = Context::new();
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
let t = fit.t_statistics(&ctx).unwrap();
close(ev(&t[0]), 0.255_152_695_240_224_77, 1e-12);
close(ev(&t[1]), 5.857_173_104_197_21, 1e-12);
close(ev(&t[2]), 10.640_157_550_951_814, 1e-12);
let p = fit.p_values(&ctx).unwrap();
close(ev(&p[0]), 0.808_768_137_709_466_2, 1e-12);
close(ev(&p[1]), 0.002_055_751_582_282_303, 1e-12);
close(ev(&p[2]), 0.000_126_857_669_237_960_7, 1e-12);
let tests = fit.coefficient_tests(&ctx).unwrap();
assert_eq!(tests.len(), 3);
close(
tests[2].statistic_f64().unwrap(),
10.640_157_550_951_814,
1e-12,
);
close(
tests[2].p_value_f64().unwrap(),
0.000_126_857_669_237_960_7,
1e-12,
);
}
#[test]
fn ols_two_regressors_conf_int_95_and_99() {
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, -1.697_438_922_482_793_7, 1e-9);
close(ci[0].upper, 2.071_543_323_711_033, 1e-9);
close(ci[1].lower, 0.389_368_432_709_537_1, 1e-9);
close(ci[1].upper, 0.998_451_423_994_658_9, 1e-9);
close(ci[2].lower, 0.807_467_270_514_176_4, 1e-9);
close(ci[2].upper, 1.321_908_369_199_232_6, 1e-9);
let ci = fit.conf_int(0.99).unwrap();
close(ci[0].lower, -2.768_908_023_731_902, 1e-9);
close(ci[0].upper, 3.143_012_424_960_141, 1e-9);
close(ci[1].lower, 0.216_214_630_799_899_22, 1e-9);
close(ci[1].upper, 1.171_605_225_904_296_6, 1e-9);
close(ci[2].lower, 0.661_218_839_068_173_2, 1e-9);
close(ci[2].upper, 1.468_156_800_645_235_7, 1e-9);
}
#[test]
fn ols_two_regressors_predict_and_intervals() {
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
let x0 = [qi(9), qi(6)];
assert_eq!(fit.predict(&x0).unwrap(), q(25051, 1954));
let mean_ci = fit.confidence_interval_mean_response(&x0, 0.95).unwrap();
close(mean_ci.lower, 11.285_745_794_051_824, 1e-9);
close(mean_ci.upper, 14.354_991_155_794_632, 1e-9);
let obs_ci = fit.prediction_interval(&x0, 0.95).unwrap();
close(obs_ci.lower, 10.329_841_215_158_392, 1e-9);
close(obs_ci.upper, 15.310_895_734_688_064, 1e-9);
assert_eq!(fit.predict(&x[2]).unwrap(), fit.fitted[2]);
}
#[test]
fn ols_two_regressors_leverage_cooks_durbin_watson_exact() {
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
assert_eq!(
fit.leverage(),
qs(&[
(2064, 4885),
(1231, 3908),
(1799, 4885),
(1487, 3908),
(996, 4885),
(6659, 19540),
(2549, 4885),
(8739, 19540)
])
);
let cooks = fit.cooks_distance().unwrap();
assert_eq!(
cooks,
qs(&[
(319040720, 10528488243),
(1323325, 64496961),
(18957966085, 58047479469),
(425616575, 1000320417747),
(2564781340, 17559336681),
(3438926314855, 28317096117387),
(18403780, 101053827),
(1842896288095, 2212253939763)
])
);
close(
data::to_f64(&[cooks[7].clone()])[0],
0.833_040_120_291_267_5,
1e-15,
);
close(
data::to_f64(&[cooks[7].clone()])[0],
0.833_040_120_291_253_7,
1e-13,
);
assert_eq!(fit.durbin_watson().unwrap(), q(1786950583, 1111611060));
close(
data::to_f64(&[fit.durbin_watson().unwrap()])[0],
1.607_532_209_152_362_7,
1e-15,
);
}
#[test]
fn ols_two_regressors_aic_bic() {
let ctx = Context::new();
let (x, y) = d2();
let fit = ols(&y, &x, true).unwrap();
close(
ev(&fit.log_likelihood(&ctx).unwrap()),
-7.308_295_515_765_012,
1e-12,
);
close(ev(&fit.aic(&ctx).unwrap()), 20.616_591_031_530_024, 1e-12);
close(ev(&fit.bic(&ctx).unwrap()), 20.854_915_656_569_53, 1e-12);
}
#[test]
fn design_builder_matches_ols_rows() {
let (x, y) = d2();
let x1: Vec<Q> = x.iter().map(|r| r[0].clone()).collect();
let x2: Vec<Q> = x.iter().map(|r| r[1].clone()).collect();
let design = Design::new().intercept().column(&x1).column(&x2);
assert!(design.has_intercept());
assert_eq!(design.n_columns(), 2);
assert_eq!(design.rows(8).unwrap(), x);
assert_eq!(design.fit(&y).unwrap(), ols(&y, &x, true).unwrap());
assert!(Design::new().intercept().column(&x1[..5]).fit(&y).is_err());
let only = Design::new().intercept().fit(&y).unwrap();
assert_eq!(only.coefficients, vec![qi(17) / qi(2)]);
assert_eq!(only.r_squared, Q::zero());
assert_eq!(only.df_model, 0);
assert!(matches!(
only.f_statistic(),
Err(SymplexError::InvalidArgument { .. })
));
assert!(only.anova_table().is_err());
}
#[test]
fn vif_two_regressors_exact() {
let (x, _) = d2();
let v = vif(&x).unwrap();
assert_eq!(v, vec![q(9891, 9770), q(9891, 9770)]);
close(data::to_f64(&v)[0], 1.012_384_851_586_489, 1e-15);
}
#[test]
fn ols_no_intercept_uses_uncentered_tss() {
let (x, y) = d3();
let fit = ols(&y, &x, false).unwrap();
assert!(!fit.has_constant());
assert_eq!(fit.coefficients, vec![q(116, 55)]);
assert_eq!(
fit.fitted,
qs(&[(116, 55), (232, 55), (348, 55), (464, 55), (116, 11)])
);
assert_eq!(
fit.residuals,
qs(&[(-6, 55), (-67, 55), (37, 55), (-24, 55), (5, 11)])
);
assert_eq!(fit.ssr, q(129, 55));
assert_eq!(fit.tss, qi(247));
assert_eq!(fit.ess, q(13456, 55));
assert_eq!(fit.r_squared, q(13456, 13585));
assert_eq!(fit.adjusted_r_squared, q(10739, 10868));
assert_eq!(fit.mse_resid, q(129, 220));
assert_eq!((fit.df_model, fit.df_resid), (1, 4));
assert_eq!(fit.f_statistic().unwrap(), q(53824, 129));
assert_eq!(fit.cov_params.to_rows(), vec![vec![q(129, 12100)]]);
}
#[test]
fn ols_no_intercept_inference_and_information_criteria() {
let ctx = Context::new();
let (x, y) = d3();
let fit = ols(&y, &x, false).unwrap();
close(
ev(&fit.standard_errors(&ctx)[0]),
0.103_252_879_014_550_41,
1e-12,
);
close(
ev(&fit.t_statistics(&ctx).unwrap()[0]),
20.426_461_026_754_474,
1e-12,
);
close(
ev(&fit.p_values(&ctx).unwrap()[0]),
3.392_120_339_091_609e-5,
1e-14,
);
close(
fit.f_test(&ctx).unwrap().p_value_f64().unwrap(),
3.392_120_339_091_610_5e-5,
1e-14,
);
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, 1.822_414_958_553_380_4, 1e-9);
close(ci[0].upper, 2.395_766_859_628_437_4, 1e-9);
let ci = fit.conf_int(0.90).unwrap();
close(ci[0].lower, 1.888_971_590_784_765_3, 1e-9);
close(ci[0].upper, 2.329_210_227_397_052_5, 1e-9);
close(
ev(&fit.log_likelihood(&ctx).unwrap()),
-5.202_295_932_761_115,
1e-12,
);
close(ev(&fit.aic(&ctx).unwrap()), 12.404_591_865_522_23, 1e-12);
close(ev(&fit.bic(&ctx).unwrap()), 12.014_029_777_956_33, 1e-12);
}
#[test]
fn ols_no_intercept_prediction_and_influence() {
let (x, y) = d3();
let fit = ols(&y, &x, false).unwrap();
assert_eq!(fit.predict(&[qi(6)]).unwrap(), q(696, 55));
let mean_ci = fit
.confidence_interval_mean_response(&[qi(6)], 0.95)
.unwrap();
close(mean_ci.lower, 10.934_489_751_320_282, 1e-9);
close(mean_ci.upper, 14.374_601_157_770_623, 1e-9);
let obs_ci = fit.prediction_interval(&[qi(6)], 0.95).unwrap();
close(obs_ci.lower, 9.919_831_181_333_315, 1e-9);
close(obs_ci.upper, 15.389_259_727_757_59, 1e-9);
assert_eq!(
fit.leverage(),
qs(&[(1, 55), (4, 55), (9, 55), (16, 55), (5, 11)])
);
assert_eq!(
fit.cooks_distance().unwrap(),
qs(&[
(4, 10449),
(71824, 335529),
(4107, 22747),
(4096, 21801),
(625, 1161)
])
);
assert_eq!(fit.durbin_watson().unwrap(), q(20659, 7095));
assert_eq!(fit.hat_matrix().unwrap().trace().unwrap(), qi(1));
}
#[test]
fn ols_detects_user_supplied_constant_column() {
let (x, y) = d1();
let with_ones: Vec<Vec<Q>> = x.iter().map(|v| vec![qi(1), v.clone()]).collect();
let fit = ols(&y, &with_ones, false).unwrap();
assert!(fit.has_constant());
assert_eq!(fit.coefficients, vec![q(5, 7), q(8, 7)]);
assert_eq!(fit.r_squared, q(64, 69));
assert_eq!(fit.adjusted_r_squared, q(21, 23));
assert_eq!(fit.df_model, 1);
assert_eq!(fit.predict(&[qi(1), qi(8)]).unwrap(), q(69, 7));
assert!(fit.predict(&[qi(8)]).is_err());
}
#[test]
fn ols_detects_implicit_constant_from_full_dummy_set() {
let ctx = Context::new();
let x = rows(&[
&[1, 0],
&[1, 0],
&[1, 0],
&[0, 1],
&[0, 1],
&[0, 1],
&[0, 1],
]);
let y = from_i64(&[2, 3, 4, 6, 7, 8, 9]);
let fit = ols(&y, &x, false).unwrap();
assert!(fit.has_constant());
assert_eq!(fit.coefficients, vec![qi(3), q(15, 2)]);
assert_eq!(fit.tss, q(292, 7));
assert_eq!(fit.r_squared, q(243, 292));
assert_eq!(fit.df_model, 1);
assert_eq!(fit.f_statistic().unwrap(), q(1215, 49));
close(
fit.f_test(&ctx).unwrap().p_value_f64().unwrap(),
0.004_177_335_830_133_132,
1e-12,
);
}
#[test]
fn wls_coefficients_and_weighted_sums_of_squares_exact() {
let (_, y) = d1();
let fit = wls(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), &d4_weights(), true).unwrap();
assert_eq!(fit.weights(), Some(&d4_weights()[..]));
assert_eq!(fit.coefficients, vec![q(29, 31), q(35, 31)]);
assert_eq!(
fit.fitted,
qs(&[
(64, 31),
(99, 31),
(134, 31),
(169, 31),
(204, 31),
(239, 31),
(274, 31)
])
);
assert_eq!(
fit.residuals,
qs(&[
(-2, 31),
(-6, 31),
(21, 31),
(-45, 31),
(-18, 31),
(9, 31),
(5, 31)
])
);
assert_eq!(fit.ssr, q(140, 31));
assert_eq!(fit.tss, q(770, 13));
assert_eq!(fit.ess, q(22050, 403));
assert_eq!(fit.r_squared, q(315, 341));
assert_eq!(fit.adjusted_r_squared, q(1549, 1705));
assert_eq!(fit.mse_resid, q(28, 31));
assert_eq!(fit.f_statistic().unwrap(), q(1575, 26));
assert_eq!((fit.df_model, fit.df_resid), (1, 5));
}
#[test]
fn wls_covariance_and_inference() {
let ctx = Context::new();
let (_, y) = d1();
let fit = wls(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), &d4_weights(), true).unwrap();
assert_eq!(
fit.normalized_cov_params().to_rows(),
vec![qs(&[(259, 558), (-53, 558)]), qs(&[(-53, 558), (13, 558)])]
);
assert_eq!(
fit.cov_params.to_rows(),
vec![
qs(&[(3626, 8649), (-742, 8649)]),
qs(&[(-742, 8649), (182, 8649)])
]
);
let se = fit.standard_errors(&ctx);
close(ev(&se[0]), 0.647_486_848_056_972, 1e-12);
close(ev(&se[1]), 0.145_061_694_228_301_45, 1e-12);
let t = fit.t_statistics(&ctx).unwrap();
close(ev(&t[0]), 1.444_792_081_530_326_6, 1e-12);
close(ev(&t[1]), 7.783_117_824_941_566_5, 1e-12);
let p = fit.p_values(&ctx).unwrap();
close(ev(&p[0]), 0.208_132_072_999_786_37, 1e-12);
close(ev(&p[1]), 0.000_560_575_680_278_43, 1e-12);
close(
fit.f_test(&ctx).unwrap().p_value_f64().unwrap(),
0.000_560_575_680_278_424_6,
1e-12,
);
close(
ev(&fit.residual_standard_error(&ctx)),
0.950_381_926_622_982_9,
1e-12,
);
}
#[test]
fn wls_log_likelihood_information_criteria_and_conf_int() {
let ctx = Context::new();
let (_, y) = d1();
let fit = wls(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), &d4_weights(), true).unwrap();
close(
ev(&fit.log_likelihood(&ctx).unwrap()),
-6.606_918_004_945_605,
1e-12,
);
close(ev(&fit.aic(&ctx).unwrap()), 17.213_836_009_891_21, 1e-12);
close(ev(&fit.bic(&ctx).unwrap()), 17.105_656_308_001_837, 1e-12);
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, -0.728_934_059_460_919_7, 1e-9);
close(ci[0].upper, 2.599_901_801_396_405_3, 1e-9);
close(ci[1].lower, 0.756_139_301_834_615_3, 1e-9);
close(ci[1].upper, 1.501_925_214_294_417_3, 1e-9);
}
#[test]
fn wls_leverage_cooks_prediction_and_durbin_watson() {
let (_, y) = d1();
let fit = wls(&y, &col(&[1, 2, 3, 4, 5, 6, 7]), &d4_weights(), true).unwrap();
let lev = fit.leverage();
assert_eq!(
lev,
qs(&[
(83, 279),
(11, 31),
(29, 93),
(43, 558),
(6, 31),
(91, 186),
(77, 279)
])
);
assert_eq!(lev.iter().fold(Q::zero(), |a, b| a + b), qi(2));
close(data::to_f64(&lev)[5], 0.489_247_311_827_956_94, 1e-15);
assert_eq!(
fit.cooks_distance().unwrap(),
qs(&[
(747, 537824),
(99, 2800),
(16443, 32768),
(31347, 297052),
(486, 4375),
(9477, 36100),
(2475, 326432)
])
);
let h = fit.hat_matrix().unwrap();
assert_eq!(h.diagonal(), lev);
let ycol = QMatrix::new(y.iter().map(|v| vec![v.clone()]).collect()).unwrap();
assert_eq!(h.matmul(&ycol).unwrap().col(0), fit.fitted);
assert_eq!(h.matmul(&h).unwrap(), h);
assert_eq!(fit.predict(&[qi(8)]).unwrap(), q(309, 31));
let mean_ci = fit
.confidence_interval_mean_response(&[qi(8)], 0.95)
.unwrap();
close(mean_ci.lower, 8.355_554_097_987_461, 1e-9);
close(mean_ci.upper, 11.579_929_772_980_282, 1e-9);
let obs_ci = fit.prediction_interval(&[qi(8)], 0.95).unwrap();
close(obs_ci.lower, 7.040_701_232_480_428, 1e-9);
close(obs_ci.upper, 12.894_782_638_487_316, 1e-9);
assert_eq!(fit.durbin_watson().unwrap(), q(6575, 2936));
close(
data::to_f64(&[q(6575, 2936)])[0],
2.239_441_416_893_73,
1e-14,
);
}
#[test]
fn wls_with_unit_weights_equals_ols_and_rejects_bad_weights() {
let (x, y) = d2();
let ones = vec![Q::one(); 8];
let w = wls(&y, &x, &ones, true).unwrap();
let o = ols(&y, &x, true).unwrap();
assert_eq!(w.coefficients, o.coefficients);
assert_eq!(w.cov_params, o.cov_params);
assert_eq!(w.r_squared, o.r_squared);
assert!(matches!(
wls(&y, &x, &ones[..7], true),
Err(SymplexError::InvalidArgument { .. })
));
let mut zero = ones.clone();
zero[3] = Q::zero();
assert!(matches!(
wls(&y, &x, &zero, true),
Err(SymplexError::InvalidArgument { .. })
));
let mut neg = ones;
neg[0] = qi(-1);
assert!(wls(&y, &x, &neg, true).is_err());
let x1: Vec<Q> = x.iter().map(|r| r[0].clone()).collect();
let x2: Vec<Q> = x.iter().map(|r| r[1].clone()).collect();
let wts = from_i64(&[1, 2, 1, 2, 1, 2, 1, 2]);
assert_eq!(
Design::new()
.intercept()
.column(&x1)
.column(&x2)
.fit_weighted(&y, &wts)
.unwrap(),
wls(&y, &x, &wts, true).unwrap()
);
}
#[test]
fn polyfit_exact_for_degrees_one_to_three() {
let x = from_i64(&[0, 1, 2, 3, 4, 5]);
let y = from_i64(&[1, 2, 6, 11, 19, 30]);
assert_eq!(polyfit(&x, &y, 1).unwrap(), vec![q(-20, 7), q(201, 35)]); assert_eq!(
polyfit(&x, &y, 2).unwrap(),
vec![q(15, 14), q(-3, 20), q(33, 28)]
);
assert_eq!(
polyfit(&x, &y, 3).unwrap(),
vec![q(19, 21), q(11, 18), q(16, 21), q(1, 18)]
);
let c = polyfit(&x, &y, 2).unwrap();
close(data::to_f64(&c)[2], 1.178_571_428_571_429_3, 1e-14);
}
#[test]
fn polyfit_interpolates_with_degree_plus_one_points_and_validates() {
let x = from_i64(&[0, 1, 2]);
let y = from_i64(&[1, 3, 9]);
assert_eq!(polyfit(&x, &y, 2).unwrap(), vec![qi(1), Q::zero(), qi(2)]);
assert!(matches!(
polyfit(&x, &y, 3),
Err(SymplexError::InvalidArgument { .. })
));
assert!(polyfit(&x, &y[..2], 1).is_err());
let xd = from_i64(&[1, 1, 2, 2]);
let yd = from_i64(&[1, 2, 3, 4]);
assert!(matches!(
polyfit(&xd, &yd, 2),
Err(SymplexError::InvalidArgument { .. })
));
assert_eq!(polyfit(&xd, &yd, 1).unwrap(), vec![q(-1, 2), qi(2)]);
}
#[test]
fn vif_three_regressors_exact() {
let x = rows(&[
&[1, 5, 1],
&[2, 3, 4],
&[3, 8, 2],
&[4, 1, 6],
&[5, 7, 3],
&[6, 2, 8],
&[7, 9, 5],
&[8, 4, 7],
]);
let v = vif(&x).unwrap();
assert_eq!(v, vec![q(84984, 8543), q(378213, 59801), q(117240, 8543)]);
let f = data::to_f64(&v);
close(f[0], 9.947_793_515_158_61, 1e-13);
close(f[1], 6.324_526_345_713_282_5, 1e-13);
close(f[2], 13.723_516_329_158_37, 1e-13);
assert_eq!(vif(&col(&[1, 2, 3, 5])).unwrap(), vec![Q::one()]);
let dup = rows(&[&[1, 2], &[2, 4], &[3, 6], &[5, 10]]);
assert!(matches!(
vif(&dup),
Err(SymplexError::InvalidArgument { .. })
));
assert!(vif(&[]).is_err());
}
#[test]
fn ols_rejects_bad_shapes() {
let (x, y) = d2();
assert!(matches!(
ols(&[], &[], true),
Err(SymplexError::InvalidArgument { .. })
));
assert!(matches!(
ols(&y[..7], &x, true),
Err(SymplexError::InvalidArgument { .. })
));
let mut ragged = x.clone();
ragged[3].push(qi(1));
assert!(matches!(
ols(&y, &ragged, true),
Err(SymplexError::InvalidArgument { .. })
));
let empty_rows: Vec<Vec<Q>> = vec![vec![]; 8];
assert!(matches!(
ols(&y, &empty_rows, false),
Err(SymplexError::InvalidArgument { .. })
));
assert_eq!(
ols(&y, &empty_rows, true).unwrap().coefficients,
vec![q(17, 2)]
);
}
#[test]
fn ols_rejects_n_le_p_and_rank_deficiency() {
let (x, y) = d2();
assert!(matches!(
ols(&y[..3], &x[..3], true),
Err(SymplexError::InvalidArgument { .. })
));
let dup: Vec<Vec<Q>> = x
.iter()
.map(|r| vec![r[0].clone(), r[0].clone() * qi(2)])
.collect();
match ols(&y, &dup, true) {
Err(SymplexError::InvalidArgument { reason, .. }) => {
assert!(reason.contains("rank deficient"), "{reason}");
}
other => panic!("expected InvalidArgument, got {other:?}"),
}
let konst: Vec<Vec<Q>> = (0..8).map(|_| vec![qi(3)]).collect();
assert!(ols(&y, &konst, true).is_err());
let flat = vec![qi(4); 8];
assert!(matches!(
ols(&flat, &x, true),
Err(SymplexError::InvalidArgument { .. })
));
}
#[test]
fn ols_perfect_fit_has_zero_residual_variance() {
let ctx = Context::new();
let x = col(&[1, 2, 3, 4, 5]);
let y = from_i64(&[3, 5, 7, 9, 11]);
let fit = ols(&y, &x, true).unwrap();
assert_eq!(fit.coefficients, vec![qi(1), qi(2)]);
assert!(fit.ssr.is_zero());
assert_eq!(fit.r_squared, Q::one());
assert_eq!(fit.adjusted_r_squared, Q::one());
assert!(fit.mse_resid.is_zero());
assert!(fit.cov_params.is_zero());
assert_eq!(fit.standard_errors(&ctx), vec![ctx.zero(), ctx.zero()]);
assert!(matches!(
fit.t_statistics(&ctx),
Err(SymplexError::InvalidArgument { .. })
));
assert!(fit.p_values(&ctx).is_err());
assert!(fit.f_statistic().is_err());
assert!(fit.cooks_distance().is_err());
assert!(fit.durbin_watson().is_err());
assert!(fit.log_likelihood(&ctx).is_err());
assert!(fit.aic(&ctx).is_err());
assert_eq!(fit.leverage().iter().fold(Q::zero(), |a, b| a + b), qi(2));
assert_eq!(fit.predict(&[qi(10)]).unwrap(), qi(21));
}
fn l1() -> (Vec<u8>, Vec<Vec<f64>>) {
let x = vec![
vec![2.0, 1.0],
vec![3.0, 4.0],
vec![1.0, 2.0],
vec![5.0, 3.0],
vec![4.0, 6.0],
vec![6.0, 2.0],
vec![2.5, 5.0],
vec![7.0, 4.0],
vec![3.5, 1.5],
vec![1.5, 6.5],
vec![5.5, 5.5],
vec![4.5, 2.5],
vec![6.5, 6.0],
vec![0.5, 3.0],
vec![3.0, 3.0],
vec![5.0, 1.0],
];
let y = vec![0u8, 0, 0, 1, 1, 1, 0, 1, 0, 0, 1, 0, 1, 0, 1, 1];
(y, x)
}
fn l2() -> (Vec<bool>, Vec<Vec<f64>>) {
let y: Vec<bool> = (0..20).map(|i| matches!(i, 7..=9 | 13..=19)).collect();
let x: Vec<Vec<f64>> = (0..20)
.map(|i| vec![if i < 10 { 0.0 } else { 1.0 }])
.collect();
(y, x)
}
fn l3() -> (Vec<u8>, Vec<Vec<f64>>) {
let x: Vec<Vec<f64>> = (1..=12).map(|i| vec![f64::from(i)]).collect();
let y = vec![0u8, 0, 0, 0, 0, 1, 0, 1, 1, 1, 1, 1];
(y, x)
}
#[test]
fn logit_two_regressors_matches_statsmodels() {
let (y, x) = l1();
let fit = logit(&y, &x, true, &LogitOpts::default()).unwrap();
assert!(fit.converged);
assert!(fit.iterations <= 20);
assert_eq!((fit.nobs, fit.df_model, fit.df_resid), (16, 2, 13));
close(fit.coefficients[0], -8.542_419_056_821_446, 1e-6);
close(fit.coefficients[1], 1.857_755_547_762_101, 1e-6);
close(fit.coefficients[2], 0.449_996_209_830_459, 1e-6);
close(fit.standard_errors[0], 4.782_605_093_317_698, 1e-6);
close(fit.standard_errors[1], 0.924_900_871_397_525_4, 1e-6);
close(fit.standard_errors[2], 0.580_841_597_355_644_8, 1e-6);
close(fit.z_values[0], -1.786_143_511_777_085, 1e-6);
close(fit.z_values[1], 2.008_599_629_660_887, 1e-6);
close(fit.z_values[2], 0.774_731_375_781_493_5, 1e-6);
close(fit.p_values[0], 0.074_076_024_534_383_36, 1e-6);
close(fit.p_values[1], 0.044_579_610_630_802_09, 1e-6);
close(fit.p_values[2], 0.438_498_406_893_896_3, 1e-6);
close(fit.log_likelihood, -4.435_039_672_866_281, 1e-6);
close(fit.null_log_likelihood, -11.090_354_888_959_125, 1e-6);
close(fit.pseudo_r_squared, 0.600_099_391_113_125_4, 1e-6);
close(fit.llr(), 13.310_630_432_185_688, 1e-6);
close(fit.aic(), 14.870_079_345_732_561, 1e-6);
close(fit.bic(), 17.187_845_512_451_904, 1e-6);
}
#[test]
fn logit_two_regressors_predictions_deviance_and_covariance() {
let (y, x) = l1();
let fit = logit(&y, &x, true, &LogitOpts::default()).unwrap();
let expected = [
0.012_408_201_943_958_73,
0.237_005_482_104_735_07,
0.003_064_898_902_249_18,
0.890_547_649_575_753_3,
0.830_416_280_586_537_3,
];
for (p, e) in fit.fitted_probabilities.iter().zip(expected) {
close(*p, e, 1e-6);
}
close(
fit.predict_proba(&[4.0, 4.0]).unwrap(),
0.665_652_740_859_788_8,
1e-6,
);
close(
fit.predict_proba(&[1.0, 1.0]).unwrap(),
0.001_956_446_108_237_788_2,
1e-6,
);
close(
fit.predict_log_odds(&[4.0, 4.0]).unwrap(),
(0.665_652_740_859_788_8_f64 / (1.0 - 0.665_652_740_859_788_8)).ln(),
1e-6,
);
close(fit.deviance, 8.870_079_345_732_563, 1e-6);
close(fit.deviance, -2.0 * fit.log_likelihood, 1e-12);
for j in 0..3 {
close(fit.cov_params[j][j], fit.standard_errors[j].powi(2), 1e-9);
}
assert!(fit.predict_proba(&[4.0]).is_err());
assert!(fit.predict_proba(&[4.0, f64::NAN]).is_err());
}
#[test]
fn logit_conf_int_and_odds_ratios() {
let (y, x) = l1();
let fit = logit(&y, &x, true, &LogitOpts::default()).unwrap();
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, -17.916_152_792_001_96, 1e-6);
close(ci[0].upper, 0.831_314_678_359_065_7, 1e-6);
close(ci[1].lower, 0.044_983_150_553_238_98, 1e-6);
close(ci[1].upper, 3.670_527_944_970_963, 1e-6);
close(ci[2].lower, -0.688_432_401_709_320_2, 1e-6);
close(ci[2].upper, 1.588_424_821_370_238_3, 1e-6);
let or = fit.odds_ratios();
close(or[0], 1.950_179_296_271_563e-4, 1e-10);
close(or[1], 6.409_335_168_957_027, 1e-6);
close(or[2], 1.568_306_241_332_357_4, 1e-6);
assert!(fit.conf_int(0.0).is_err());
assert!(fit.conf_int(1.0).is_err());
}
#[test]
fn logit_binary_regressor_has_closed_form() {
let (y, x) = l2();
let fit = logit(&y, &x, true, &LogitOpts::default()).unwrap();
assert!(fit.converged);
close(fit.coefficients[0], (3.0_f64 / 7.0).ln(), 1e-9);
close(fit.coefficients[1], (49.0_f64 / 9.0).ln(), 1e-9);
close(fit.standard_errors[0], 0.690_065_559_342_354_3, 1e-6);
close(fit.standard_errors[1], 0.975_900_072_948_533_2, 1e-6);
close(fit.z_values[0], -1.227_851_251_111_118_8, 1e-6);
close(fit.z_values[1], 1.736_443_891_898_117, 1e-6);
close(fit.p_values[0], 0.219_502_812_283_000_73, 1e-6);
close(fit.p_values[1], 0.082_485_377_115_864_68, 1e-6);
close(fit.log_likelihood, -12.217_286_041_097_868, 1e-6);
close(fit.null_log_likelihood, -13.862_943_611_198_906, 1e-6);
close(fit.null_log_likelihood, 20.0 * 0.5_f64.ln(), 1e-12);
close(fit.pseudo_r_squared, 0.118_709_100_769_307_52, 1e-6);
close(fit.llr(), 3.291_315_140_202_076_6, 1e-6);
close(fit.aic(), 28.434_572_082_195_736, 1e-6);
close(fit.bic(), 30.426_036_629_303_717, 1e-6);
let ci = fit.conf_int(0.95).unwrap();
close(ci[0].lower, -2.199_801_503_669_705_4, 1e-6);
close(ci[0].upper, 0.505_205_782_895_298, 1e-6);
close(ci[1].lower, -0.218_133_274_714_729_1, 1e-6);
close(ci[1].upper, 3.607_324_716_263_544, 1e-6);
close(fit.predict_proba(&[0.0]).unwrap(), 0.3, 1e-9);
close(fit.predict_proba(&[1.0]).unwrap(), 0.7, 1e-9);
let or = fit.odds_ratios();
close(or[0], 3.0 / 7.0, 1e-9);
close(or[1], 49.0 / 9.0, 1e-9);
}
#[test]
fn logit_near_separated_data_still_fits() {
let (y, x) = l3();
let fit = logit(&y, &x, true, &LogitOpts::default()).unwrap();
assert!(fit.converged);
close(fit.coefficients[0], -8.498_852_464_660_226, 1e-6);
close(fit.coefficients[1], 1.307_515_763_793_881, 1e-6);
close(fit.standard_errors[0], 5.526_614_625_575_412, 1e-6);
close(fit.standard_errors[1], 0.831_836_584_406_527_1, 1e-6);
close(fit.z_values[0], -1.537_804_431_908_504, 1e-6);
close(fit.z_values[1], 1.571_842_100_124_421_4, 1e-6);
close(fit.p_values[0], 0.124_096_439_645_944_93, 1e-6);
close(fit.p_values[1], 0.115_987_175_147_813_05, 1e-6);
close(fit.log_likelihood, -2.510_538_677_467_477_6, 1e-6);
close(fit.null_log_likelihood, -8.317_766_166_719_343, 1e-6);
close(fit.pseudo_r_squared, 0.698_171_525_004_811, 1e-6);
close(fit.aic(), 9.021_077_354_934_956, 1e-6);
close(fit.bic(), 9.990_890_654_510_956, 1e-6);
let expected = [
0.000_752_515_096_878_98,
0.002_776_397_109_706_61,
0.010_187_992_897_796_41,
0.036_657_554_828_494_33,
0.123_329_275_979_167_24,
];
for (p, e) in fit.fitted_probabilities.iter().zip(expected) {
close(*p, e, 1e-6);
}
close(fit.predict_proba(&[6.5]).unwrap(), 0.5, 1e-6);
}
#[test]
fn logit_without_intercept_matches_statsmodels() {
let (y, x) = l1();
let centered: Vec<Vec<f64>> = x.iter().map(|r| vec![r[0] - 4.0, r[1] - 3.5]).collect();
let fit = logit(&y, ¢ered, false, &LogitOpts::default()).unwrap();
assert!(fit.converged);
assert_eq!(fit.n_params(), 2);
assert_eq!((fit.df_model, fit.df_resid), (1, 14));
close(fit.coefficients[0], 1.777_350_889_922_984_4, 1e-6);
close(fit.coefficients[1], 0.389_543_006_851_446_1, 1e-6);
close(fit.standard_errors[0], 0.880_840_547_652_612_4, 1e-6);
close(fit.standard_errors[1], 0.526_530_807_738_018_2, 1e-6);
close(fit.z_values[0], 2.017_789_592_746_971_6, 1e-6);
close(fit.z_values[1], 0.739_829_467_006_739_7, 1e-6);
close(fit.p_values[0], 0.043_613_179_231_309_94, 1e-6);
close(fit.p_values[1], 0.459_403_476_623_750_14, 1e-6);
close(fit.log_likelihood, -4.572_400_264_628_698_5, 1e-6);
close(fit.null_log_likelihood, -11.090_354_888_959_125, 1e-6);
close(fit.pseudo_r_squared, 0.587_713_800_828_799_5, 1e-6);
close(fit.aic(), 13.144_800_529_257_397, 1e-6);
close(fit.bic(), 14.689_977_973_736_958, 1e-6);
close(fit.predict_proba(&[0.0, 0.0]).unwrap(), 0.5, 1e-12);
close(
fit.predict_proba(&[1.0, -1.0]).unwrap(),
0.800_242_053_560_043_6,
1e-6,
);
let ci = fit.conf_int(0.90).unwrap();
close(ci[0].lower, 0.328_497_120_350_663_9, 1e-6);
close(ci[0].upper, 3.226_204_659_495_304_7, 1e-6);
close(ci[1].lower, -0.476_523_101_958_121_25, 1e-6);
close(ci[1].upper, 1.255_609_115_661_013_4, 1e-6);
}
#[test]
fn logit_accepts_bool_u8_i64_and_f64_outcomes() {
let (y_u8, x) = l1();
let y_bool: Vec<bool> = y_u8.iter().map(|&v| v == 1).collect();
let y_i64: Vec<i64> = y_u8.iter().map(|&v| i64::from(v)).collect();
let y_f64: Vec<f64> = y_u8.iter().map(|&v| f64::from(v)).collect();
let opts = LogitOpts::default();
let a = logit(&y_u8, &x, true, &opts).unwrap();
assert_eq!(logit(&y_bool, &x, true, &opts).unwrap(), a);
assert_eq!(logit(&y_i64, &x, true, &opts).unwrap(), a);
assert_eq!(logit(&y_f64, &x, true, &opts).unwrap(), a);
let mut bad = y_u8.clone();
bad[0] = 2;
assert!(matches!(
logit(&bad, &x, true, &opts),
Err(SymplexError::InvalidArgument { .. })
));
let mut badf = y_f64.clone();
badf[0] = 0.5;
assert!(logit(&badf, &x, true, &opts).is_err());
}
#[test]
fn logit_detects_complete_separation() {
let x: Vec<Vec<f64>> = (1..=8).map(|i| vec![f64::from(i)]).collect();
let y = vec![0u8, 0, 0, 0, 1, 1, 1, 1];
match logit(&y, &x, true, &LogitOpts::default()) {
Err(SymplexError::ComputationFailed { reason, .. }) => {
assert!(reason.contains("separation"), "{reason}");
}
other => panic!("expected ComputationFailed, got {other:?}"),
}
let opts = LogitOpts {
max_iter: 2000,
tol: 1e-12,
};
assert!(matches!(
logit(&y, &x, true, &opts),
Err(SymplexError::ComputationFailed { .. })
));
}
#[test]
fn logit_detects_quasi_complete_separation() {
let x: Vec<Vec<f64>> = (0..14)
.map(|i| vec![if i < 8 { 0.0 } else { 1.0 }])
.collect();
let y = vec![0u8, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1];
match logit(&y, &x, true, &LogitOpts::default()) {
Err(SymplexError::ComputationFailed { reason, .. }) => {
assert!(reason.contains("separation"), "{reason}");
}
other => panic!("expected ComputationFailed, got {other:?}"),
}
assert!(matches!(
logit(
&y,
&x,
true,
&LogitOpts {
max_iter: 5000,
tol: 1e-10
}
),
Err(SymplexError::ComputationFailed { .. })
));
}
#[test]
fn logit_rejects_invalid_input() {
let (y, x) = l1();
let opts = LogitOpts::default();
assert!(matches!(
logit::<u8>(&[], &[], true, &opts),
Err(SymplexError::InvalidArgument { .. })
));
assert!(matches!(
logit(&[1u8; 16], &x, true, &opts),
Err(SymplexError::InvalidArgument { .. })
));
assert!(matches!(
logit(&y[..15], &x, true, &opts),
Err(SymplexError::InvalidArgument { .. })
));
let mut ragged = x.clone();
ragged[2].push(1.0);
assert!(logit(&y, &ragged, true, &opts).is_err());
let mut nan = x.clone();
nan[2][0] = f64::NAN;
assert!(logit(&y, &nan, true, &opts).is_err());
assert!(matches!(
logit(&y[..3], &x[..3], true, &opts),
Err(SymplexError::InvalidArgument { .. })
));
let dup: Vec<Vec<f64>> = x.iter().map(|r| vec![r[0], 2.0 * r[0]]).collect();
match logit(&y, &dup, true, &opts) {
Err(SymplexError::InvalidArgument { reason, .. }) => {
assert!(reason.contains("rank deficient"), "{reason}");
}
other => panic!("expected InvalidArgument, got {other:?}"),
}
let empty_rows: Vec<Vec<f64>> = vec![vec![]; 16];
assert!(logit(&y, &empty_rows, false, &opts).is_err());
assert!(
logit(
&y,
&x,
true,
&LogitOpts {
max_iter: 0,
tol: 1e-8
}
)
.is_err()
);
assert!(
logit(
&y,
&x,
true,
&LogitOpts {
max_iter: 10,
tol: 0.0
}
)
.is_err()
);
}
#[test]
fn logit_reports_non_convergence_when_iteration_budget_is_tiny() {
let (y, x) = l1();
let fit = logit(
&y,
&x,
true,
&LogitOpts {
max_iter: 1,
tol: 1e-10,
},
)
.unwrap();
assert!(!fit.converged);
assert_eq!(fit.iterations, 1);
let full = logit(&y, &x, true, &LogitOpts::default()).unwrap();
assert!(full.log_likelihood > fit.log_likelihood);
}