#![allow(unsafe_code)]
use crate::distributions::NormalDistribution;
const LN2: f64 = std::f64::consts::LN_2;
const LOG2_E: f64 = std::f64::consts::LOG2_E;
const INV_SQRT_2PI: f64 = 0.398_942_280_401_432_7;
const EXP_TAYLOR: [f64; 10] = [
1.0 / 39_916_800.0, 1.0 / 3_628_800.0, 1.0 / 362_880.0, 1.0 / 40_320.0, 1.0 / 5_040.0, 1.0 / 720.0, 1.0 / 120.0, 1.0 / 24.0, 1.0 / 6.0, 1.0 / 2.0, ];
const EXP_BIAS: i64 = 1023;
const EXP_SHIFT: i32 = 52;
fn simd_exp_scalar(x: f64) -> f64 {
x.exp()
}
pub(super) fn normal_pdf_into(dist: &NormalDistribution, xs: &[f64], out: &mut [f64]) {
let inv_sigma = 1.0 / dist.standard_deviation;
let norm = INV_SQRT_2PI * inv_sigma;
let mean = dist.mean;
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe {
pdf_neon(mean, inv_sigma, norm, xs, out);
}
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
unsafe {
pdf_avx2(mean, inv_sigma, norm, xs, out);
}
return;
}
}
pdf_scalar(mean, inv_sigma, norm, xs, out);
}
fn pdf_scalar(mean: f64, inv_sigma: f64, norm: f64, xs: &[f64], out: &mut [f64]) {
for (o, &x) in out.iter_mut().zip(xs) {
let z = (x - mean) * inv_sigma;
*o = simd_exp_scalar(-0.5 * z * z) * norm;
}
}
pub(super) fn normal_cdf_into(dist: &NormalDistribution, xs: &[f64], out: &mut [f64]) {
let inv_scale = 1.0 / (dist.standard_deviation * std::f64::consts::SQRT_2);
let mean = dist.mean;
for (o, &x) in out.iter_mut().zip(xs) {
let z = (x - mean) * inv_scale;
*o = 0.5 * (1.0 + crate::special::erf(z));
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn pdf_neon(mean: f64, inv_sigma: f64, norm: f64, xs: &[f64], out: &mut [f64]) {
use std::arch::aarch64::{vdupq_n_f64, vld1q_f64, vmulq_f64, vst1q_f64, vsubq_f64};
let lanes = 2;
let n = xs.len().min(out.len());
let body = n - (n % lanes);
let vmean = vdupq_n_f64(mean);
let vinv = vdupq_n_f64(inv_sigma);
let vnorm = vdupq_n_f64(norm);
let vneg_half = vdupq_n_f64(-0.5);
let mut i = 0;
while i < body {
let vx = vld1q_f64(xs.as_ptr().add(i));
let z = vmulq_f64(vsubq_f64(vx, vmean), vinv);
let arg = vmulq_f64(vmulq_f64(z, z), vneg_half);
let e = exp_neon(arg);
let res = vmulq_f64(e, vnorm);
vst1q_f64(out.as_mut_ptr().add(i), res);
i += lanes;
}
if body < n {
pdf_scalar(
mean,
inv_sigma,
norm,
xs.get(body..n).unwrap_or(&[]),
out.get_mut(body..n).unwrap_or(&mut []),
);
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn exp_neon(x: std::arch::aarch64::float64x2_t) -> std::arch::aarch64::float64x2_t {
use std::arch::aarch64::{
vaddq_f64, vaddq_s64, vcvtnq_s64_f64, vcvtq_f64_s64, vdupq_n_f64, vdupq_n_s64, vfmaq_f64,
vmulq_f64, vreinterpretq_f64_s64, vshlq_n_s64, vsubq_f64,
};
let log2e = vdupq_n_f64(LOG2_E);
let neg_ln2 = vdupq_n_f64(-LN2);
let ki = vcvtnq_s64_f64(vmulq_f64(x, log2e));
let kf = vcvtq_f64_s64(ki);
let r = vfmaq_f64(x, kf, neg_ln2);
let mut p = vdupq_n_f64(EXP_TAYLOR[0]);
for &c in &EXP_TAYLOR[1..] {
p = vfmaq_f64(vdupq_n_f64(c), p, r);
}
let r2 = vmulq_f64(r, r);
let _ = vsubq_f64; let exp_r = vaddq_f64(vfmaq_f64(r, r2, p), vdupq_n_f64(1.0));
let biased = vaddq_s64(ki, vdupq_n_s64(EXP_BIAS));
let pow2 = vreinterpretq_f64_s64(vshlq_n_s64::<EXP_SHIFT>(biased));
vmulq_f64(exp_r, pow2)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn pdf_avx2(mean: f64, inv_sigma: f64, norm: f64, xs: &[f64], out: &mut [f64]) {
use std::arch::x86_64::{
_mm256_loadu_pd, _mm256_mul_pd, _mm256_set1_pd, _mm256_storeu_pd, _mm256_sub_pd,
};
let lanes = 4;
let n = xs.len().min(out.len());
let body = n - (n % lanes);
let vmean = _mm256_set1_pd(mean);
let vinv = _mm256_set1_pd(inv_sigma);
let vnorm = _mm256_set1_pd(norm);
let vneg_half = _mm256_set1_pd(-0.5);
let mut i = 0;
while i < body {
unsafe {
let vx = _mm256_loadu_pd(xs.as_ptr().add(i));
let z = _mm256_mul_pd(_mm256_sub_pd(vx, vmean), vinv);
let arg = _mm256_mul_pd(_mm256_mul_pd(z, z), vneg_half);
let e = exp_avx2(arg);
let res = _mm256_mul_pd(e, vnorm);
_mm256_storeu_pd(out.as_mut_ptr().add(i), res);
}
i += lanes;
}
if body < n {
pdf_scalar(
mean,
inv_sigma,
norm,
xs.get(body..n).unwrap_or(&[]),
out.get_mut(body..n).unwrap_or(&mut []),
);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn exp_avx2(x: std::arch::x86_64::__m256d) -> std::arch::x86_64::__m256d {
use std::arch::x86_64::{
_mm256_add_epi64, _mm256_add_pd, _mm256_castsi256_pd, _mm256_cvtepi32_epi64,
_mm256_cvtpd_epi32, _mm256_fmadd_pd, _mm256_mul_pd, _mm256_round_pd, _mm256_set1_epi64x,
_mm256_set1_pd, _mm256_slli_epi64,
};
const ROUND_NEAREST: i32 = 0x00;
let log2e = _mm256_set1_pd(LOG2_E);
let neg_ln2 = _mm256_set1_pd(-LN2);
let kf = _mm256_round_pd::<ROUND_NEAREST>(_mm256_mul_pd(x, log2e));
let r = _mm256_fmadd_pd(kf, neg_ln2, x);
let mut p = _mm256_set1_pd(EXP_TAYLOR[0]);
for &c in &EXP_TAYLOR[1..] {
p = _mm256_fmadd_pd(p, r, _mm256_set1_pd(c));
}
let r2 = _mm256_mul_pd(r, r);
let exp_r = _mm256_add_pd(_mm256_fmadd_pd(r2, p, r), _mm256_set1_pd(1.0));
let ki32 = _mm256_cvtpd_epi32(kf);
let ki64 = _mm256_cvtepi32_epi64(ki32);
let biased = _mm256_add_epi64(ki64, _mm256_set1_epi64x(EXP_BIAS));
let pow2 = _mm256_castsi256_pd(_mm256_slli_epi64::<EXP_SHIFT>(biased));
_mm256_mul_pd(exp_r, pow2)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::distributions::Pdf;
#[test]
fn batch_pdf_matches_stdlib_gaussian() {
let dist = NormalDistribution {
mean: 0.0,
standard_deviation: 1.0,
..Default::default()
};
let xs: Vec<f64> = (0..251).map(|i| f64::from(i).mul_add(0.04, -5.0)).collect();
let mut out = vec![0.0; xs.len()];
normal_pdf_into(&dist, &xs, &mut out);
let inv_sqrt_2pi = 1.0 / (2.0 * std::f64::consts::PI).sqrt();
for (o, &x) in out.iter().zip(&xs) {
let want = inv_sqrt_2pi * (-0.5 * x * x).exp();
let rel = ((o - want) / want).abs();
assert!(
rel < 1e-12,
"pdf({x}) rel error {rel}: got {o}, want {want}"
);
}
}
#[test]
fn batch_pdf_matches_scalar_pdf() {
let dist = NormalDistribution {
mean: 0.3,
standard_deviation: 1.7,
..Default::default()
};
let xs: Vec<f64> = (0..103).map(|i| f64::from(i).mul_add(0.1, -5.0)).collect();
let mut out = vec![0.0; xs.len()];
normal_pdf_into(&dist, &xs, &mut out);
for (o, &x) in out.iter().zip(&xs) {
let want = dist.pdf(x);
assert!(
(o - want).abs() < 1e-12,
"batch pdf {o} != scalar pdf {want} at x={x}"
);
}
}
#[test]
fn batch_cdf_matches_scalar_cdf() {
use crate::distributions::Cdf;
let dist = NormalDistribution {
mean: -0.5,
standard_deviation: 2.0,
..Default::default()
};
let xs: Vec<f64> = (0..97).map(|i| f64::from(i).mul_add(0.1, -5.0)).collect();
let mut out = vec![0.0; xs.len()];
normal_cdf_into(&dist, &xs, &mut out);
for (o, &x) in out.iter().zip(&xs) {
let want = dist.cdf(x);
assert!(
(o - want).abs() < 1e-12,
"batch cdf {o} != scalar cdf {want} at x={x}"
);
}
}
}