use crate::spec::{BinomialLink, Family};
mod sealed {
pub trait Sealed {}
impl Sealed for f64 {}
impl<const N: usize> Sealed for crate::dual::Dual<N> {}
impl<const N: usize, const H: usize> Sealed for crate::dual::HyperDual<N, H> {}
}
#[doc(hidden)]
pub trait Scalar:
sealed::Sealed
+ Copy
+ Send
+ Sync
+ std::fmt::Debug
+ std::ops::Add<Output = Self>
+ std::ops::Sub<Output = Self>
+ std::ops::Mul<Output = Self>
+ std::ops::Div<Output = Self>
+ std::ops::Neg<Output = Self>
+ std::ops::AddAssign
+ std::ops::SubAssign
+ std::ops::MulAssign
+ std::ops::DivAssign
{
const ZERO: Self;
const ONE: Self;
const IS_F64: bool;
fn from_f64(v: f64) -> Self;
fn value(self) -> f64;
fn abs(self) -> Self;
fn sqrt(self) -> Self;
fn max_f64(self, other: f64) -> Self;
fn clamp_f64(self, lo: f64, hi: f64) -> Self;
fn mul_add(self, a: Self, b: Self) -> Self;
fn exp(self) -> Self;
fn exp_m1(self) -> Self;
fn ln(self) -> Self;
fn log1pexp(self) -> Self;
fn sigmoid(self) -> Self;
fn probit_cdf(self) -> Self;
fn ln_gamma(self) -> Self;
#[allow(clippy::too_many_arguments)]
fn family_pass(
family: Family,
nb_theta: f64,
eta: &mut [Self],
y: &[f64],
prior_w: &[f64],
weighted: bool,
yeta: Self,
prob: &mut [Self],
w: &mut [Self],
z: &mut [Self],
) -> (Self, bool) {
generic_family_pass(
family, nb_theta, eta, y, prior_w, weighted, yeta, prob, w, z,
)
}
fn chol_lower(a: &[Self], dim: usize, l_out: &mut [Self]) -> bool;
fn syrk_lower_sub(
bt: &[Self],
t_dim: usize,
w_tot: usize,
tail: &mut [Self],
scratch: &mut [f64],
);
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn generic_family_pass<T: Scalar>(
family: Family,
nb_theta: f64,
eta: &mut [T],
y: &[f64],
prior_w: &[f64],
weighted: bool,
yeta: T,
prob: &mut [T],
w: &mut [T],
z: &mut [T],
) -> (T, bool) {
let n = eta.len();
let clamp = crate::glm::WEIGHT_CLAMP;
if matches!(
family,
Family::Binomial {
link: BinomialLink::Logit
}
) && !weighted
{
let mut lp = T::ZERO;
for i in 0..n {
let e = eta[i];
let p = e.sigmoid();
prob[i] = p;
w[i] = (p * (T::ONE - p)).max_f64(clamp);
lp += e.log1pexp();
if !z.is_empty() {
z[i] = e + (T::from_f64(y[i]) - p) / w[i];
}
}
return (T::from_f64(2.0) * (lp - yeta), false);
}
let mut dev = T::ZERO;
let mut infeasible = false;
for i in 0..n {
infeasible |= crate::family::eta_infeasible(family, eta[i]);
eta[i] = crate::family::clamp_eta(family, eta[i]);
let (mu, w_raw, r) = crate::family::irls_weight_and_resid(family, nb_theta, y[i], eta[i]);
let pw = if prior_w.is_empty() { 1.0 } else { prior_w[i] };
prob[i] = mu;
w[i] = (T::from_f64(pw) * w_raw).max_f64(clamp);
if !z.is_empty() {
z[i] = eta[i] + r;
}
dev += T::from_f64(pw) * crate::family::dev_resid(family, nb_theta, y[i], mu);
}
(dev, infeasible)
}
impl Scalar for f64 {
const ZERO: f64 = 0.0;
const ONE: f64 = 1.0;
const IS_F64: bool = true;
#[inline(always)]
fn from_f64(v: f64) -> f64 {
v
}
#[inline(always)]
fn value(self) -> f64 {
self
}
#[inline(always)]
fn abs(self) -> f64 {
f64::abs(self)
}
#[inline(always)]
fn sqrt(self) -> f64 {
f64::sqrt(self)
}
#[inline(always)]
fn max_f64(self, other: f64) -> f64 {
f64::max(self, other)
}
#[inline(always)]
fn clamp_f64(self, lo: f64, hi: f64) -> f64 {
f64::clamp(self, lo, hi)
}
#[inline(always)]
fn mul_add(self, a: f64, b: f64) -> f64 {
f64::mul_add(self, a, b)
}
#[inline(always)]
fn exp(self) -> f64 {
f64::exp(self)
}
#[inline(always)]
fn exp_m1(self) -> f64 {
f64::exp_m1(self)
}
#[inline(always)]
fn ln(self) -> f64 {
f64::ln(self)
}
#[inline(always)]
fn log1pexp(self) -> f64 {
crate::simd_transcendental::scalar_log1pexp(self)
}
#[inline(always)]
fn sigmoid(self) -> f64 {
crate::glm::sigmoid_stable(self)
}
#[inline(always)]
fn probit_cdf(self) -> f64 {
crate::simd_transcendental::phi_hp(self)
}
#[inline(always)]
fn ln_gamma(self) -> f64 {
crate::simd_transcendental::ln_gamma(self)
}
#[inline(always)]
fn family_pass(
family: Family,
nb_theta: f64,
eta: &mut [f64],
y: &[f64],
prior_w: &[f64],
weighted: bool,
yeta: f64,
prob: &mut [f64],
w: &mut [f64],
z: &mut [f64],
) -> (f64, bool) {
crate::simd_transcendental::family_pass(
family, nb_theta, eta, y, prior_w, weighted, yeta, prob, w, z,
)
}
fn chol_lower(a: &[f64], dim: usize, l_out: &mut [f64]) -> bool {
use faer::dyn_stack::{MemBuffer, MemStack};
use faer::linalg::cholesky::llt::factor::{
cholesky_in_place, cholesky_in_place_scratch, LltRegularization,
};
use faer::{Par, Spec};
l_out[..dim * dim].fill(0.0);
for j in 0..dim {
let col = j * dim;
l_out[col + j..col + dim].copy_from_slice(&a[col + j..col + dim]);
}
let mut mem = MemBuffer::new(cholesky_in_place_scratch::<f64>(
dim,
Par::Seq,
Spec::default(),
));
let l = faer::MatMut::from_column_major_slice_mut(&mut l_out[..dim * dim], dim, dim);
cholesky_in_place(
l,
LltRegularization::default(),
Par::Seq,
MemStack::new(&mut mem),
Spec::default(),
)
.is_ok()
}
fn syrk_lower_sub(
bt: &[f64],
t_dim: usize,
w_tot: usize,
tail: &mut [f64],
_scratch: &mut [f64],
) {
tri_lower_sub_gemm(tail, t_dim, bt, bt, w_tot);
}
}
pub(crate) fn chol_lower_generic<T: Scalar>(a: &[T], dim: usize, l_out: &mut [T]) -> bool {
for j in 0..dim {
let mut d = a[j * dim + j];
for k in 0..j {
let l = l_out[k * dim + j];
d -= l * l;
}
if !(d.value().is_finite() && d.value() > 0.0) {
return false;
}
let ljj = d.sqrt();
l_out[j * dim + j] = ljj;
for i in (j + 1)..dim {
let mut s = a[j * dim + i];
for k in 0..j {
s -= l_out[k * dim + i] * l_out[k * dim + j];
}
l_out[j * dim + i] = s / ljj;
}
}
true
}
pub(crate) fn tri_lower_sub_gemm(
tail: &mut [f64],
t_dim: usize,
a: &[f64],
b: &[f64],
w_tot: usize,
) {
use faer::linalg::matmul::triangular::{matmul, BlockStructure};
let a = faer::MatRef::from_column_major_slice(&a[..t_dim * w_tot], t_dim, w_tot);
let b = faer::MatRef::from_column_major_slice(&b[..t_dim * w_tot], t_dim, w_tot);
let tail = faer::MatMut::from_column_major_slice_mut(&mut tail[..t_dim * t_dim], t_dim, t_dim);
matmul(
tail,
BlockStructure::TriangularLower,
faer::Accum::Add,
a,
BlockStructure::Rectangular,
b.transpose(),
BlockStructure::Rectangular,
-1.0,
faer::Par::Seq,
);
}
pub(crate) fn syrk_lower_sub_generic<T: Scalar>(
bt: &[T],
t_dim: usize,
w_tot: usize,
tail: &mut [T],
) {
for j in 0..t_dim {
for i in j..t_dim {
let mut s = T::ZERO;
for k in 0..w_tot {
s += bt[k * t_dim + i] * bt[k * t_dim + j];
}
tail[j * t_dim + i] -= s;
}
}
}
#[cfg(test)]
mod tests {
use super::Scalar;
#[test]
fn f64_impl_is_the_identity_binding() {
for &v in &[-3.5_f64, -1e-9, 0.0, 0.5, 1.0, 40.0, 700.0] {
assert_eq!(Scalar::exp(v), v.exp());
assert_eq!(Scalar::exp_m1(v), v.exp_m1());
assert_eq!(Scalar::sigmoid(v), crate::glm::sigmoid_stable(v));
assert_eq!(Scalar::probit_cdf(v), crate::simd_transcendental::phi_hp(v));
assert_eq!(Scalar::value(v), v);
assert_eq!(<f64 as Scalar>::from_f64(v), v);
}
for &v in &[1e-8_f64, 0.5, 1.0, 12.0] {
assert_eq!(Scalar::ln(v), v.ln());
assert_eq!(Scalar::sqrt(v), v.sqrt());
assert_eq!(Scalar::ln_gamma(v), crate::simd_transcendental::ln_gamma(v));
}
}
#[test]
fn generic_default_agrees_with_simd_family_pass() {
use crate::spec::{BinomialLink, Family, GammaLink, PoissonLink};
type PassResult = (f64, bool, Vec<f64>, Vec<f64>, Vec<f64>);
fn run(
family: Family,
eta0: &[f64],
y: &[f64],
weighted: bool,
) -> (PassResult, PassResult) {
let n = eta0.len();
let prior_w: Vec<f64> = vec![];
let yeta = if weighted {
0.0
} else {
eta0.iter().zip(y).map(|(&e, &yy)| yy * e).sum()
};
let mut eta_a = eta0.to_vec();
let mut prob_a = vec![0.0; n];
let mut w_a = vec![0.0; n];
let mut z_a = vec![0.0; n];
let (dev_a, inf_a) = super::generic_family_pass(
family,
f64::NAN,
&mut eta_a,
y,
&prior_w,
weighted,
yeta,
&mut prob_a,
&mut w_a,
&mut z_a,
);
let mut eta_b = eta0.to_vec();
let mut prob_b = vec![0.0; n];
let mut w_b = vec![0.0; n];
let mut z_b = vec![0.0; n];
let (dev_b, inf_b) = crate::simd_transcendental::family_pass(
family,
f64::NAN,
&mut eta_b,
y,
&prior_w,
weighted,
yeta,
&mut prob_b,
&mut w_b,
&mut z_b,
);
(
(dev_a, inf_a, eta_a, prob_a, w_a),
(dev_b, inf_b, eta_b, prob_b, w_b),
)
}
let etas: Vec<f64> = vec![-3.0, -1.0, -0.1, 0.0, 0.2, 1.0, 3.0, 8.0];
let ys: Vec<f64> = vec![0.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0];
let (a, b) = run(
Family::Binomial {
link: BinomialLink::Logit,
},
&etas,
&ys,
false,
);
assert!(
(a.0 - b.0).abs() < 1e-9,
"logit deviance drifted: {a:?} vs {b:?}"
);
for (pa, pb) in a.3.iter().zip(&b.3) {
assert!((pa - pb).abs() < 1e-12, "logit prob drifted");
}
let (a, b) = run(
Family::Binomial {
link: BinomialLink::Probit,
},
&etas,
&ys,
false,
);
assert!(
(a.0 - b.0).abs() < 1e-9,
"probit deviance drifted: {a:?} vs {b:?}"
);
for (pa, pb) in a.3.iter().zip(&b.3) {
assert!((pa - pb).abs() < 1e-12, "probit prob drifted");
}
let (a, b) = run(
Family::Poisson {
link: PoissonLink::Log,
},
&etas,
&[0.0, 1.0, 2.0, 0.0, 3.0, 1.0, 5.0, 4.0],
false,
);
assert!((a.0 - b.0).abs() < 1e-12, "poisson-log deviance drifted");
for (pa, pb) in a.3.iter().zip(&b.3) {
assert!(
(pa - pb).abs() < 1e-15,
"poisson-log mu drifted beyond 1 ULP"
);
}
let (a, b) = run(
Family::Gamma {
link: GammaLink::Log,
},
&etas,
&[0.5, 1.2, 2.0, 0.3, 4.0, 1.0, 6.0, 2.5],
false,
);
assert!((a.0 - b.0).abs() < 1e-12, "gamma-log deviance drifted");
for (pa, pb) in a.3.iter().zip(&b.3) {
assert!((pa - pb).abs() < 1e-15, "gamma-log mu drifted beyond 1 ULP");
}
}
#[test]
fn scalar_trait_is_still_sealed_and_send_sync() {
fn assert_send_sync<T: Scalar + Send + Sync>() {}
assert_send_sync::<f64>();
assert_send_sync::<crate::dual::Dual<4>>();
assert_send_sync::<crate::dual::HyperDual<4, 10>>();
}
}