#![allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::cast_sign_loss,
clippy::similar_names,
clippy::missing_errors_doc,
clippy::missing_panics_doc
)]
#![allow(clippy::missing_safety_doc)]
#![allow(clippy::wildcard_imports)]
#![allow(clippy::cast_ptr_alignment)]
#![allow(clippy::ptr_as_ptr)]
#![allow(clippy::doc_markdown)]
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SimdSupport {
None,
Avx2,
}
#[must_use]
pub fn detect_simd_support() -> SimdSupport {
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") {
SimdSupport::Avx2
} else {
SimdSupport::None
}
}
#[cfg(not(target_arch = "x86_64"))]
{
SimdSupport::None
}
}
pub const DT_OVER_TAU_MAX: i32 = 1884;
#[must_use]
pub fn dt_over_tau(dt_us: u32, tau_membrane_us: u32) -> i32 {
let exact = crate::lif_neuron::dt_over_tau(dt_us, tau_membrane_us);
i32::try_from(exact)
.unwrap_or(i32::MAX)
.min(DT_OVER_TAU_MAX)
}
pub fn integrate_lif_batch(
membrane: &mut [i16],
resting: &[i16],
input_currents: &[i16],
resistance: &[i16],
threshold: &[i16],
dt_over_tau: i32,
spikes_out: &mut [bool],
) {
let n = membrane.len();
assert_eq!(resting.len(), n, "resting.len() != membrane.len()");
assert_eq!(
input_currents.len(),
n,
"input_currents.len() != membrane.len()"
);
assert_eq!(resistance.len(), n, "resistance.len() != membrane.len()");
assert_eq!(threshold.len(), n, "threshold.len() != membrane.len()");
assert_eq!(spikes_out.len(), n, "spikes_out.len() != membrane.len()");
let dt_over_tau = dt_over_tau.clamp(-DT_OVER_TAU_MAX, DT_OVER_TAU_MAX);
#[cfg(target_arch = "x86_64")]
if matches!(detect_simd_support(), SimdSupport::Avx2) {
unsafe {
integrate_batch_avx2(
membrane,
resting,
input_currents,
resistance,
threshold,
dt_over_tau,
spikes_out,
);
}
return;
}
integrate_batch_scalar(
membrane,
resting,
input_currents,
resistance,
threshold,
dt_over_tau,
spikes_out,
);
}
pub fn integrate_batch_scalar(
membrane: &mut [i16],
resting: &[i16],
input_currents: &[i16],
resistance: &[i16],
threshold: &[i16],
dt_over_tau: i32,
spikes_out: &mut [bool],
) {
let n = membrane.len();
assert_eq!(resting.len(), n, "resting.len() != membrane.len()");
assert_eq!(
input_currents.len(),
n,
"input_currents.len() != membrane.len()"
);
assert_eq!(resistance.len(), n, "resistance.len() != membrane.len()");
assert_eq!(threshold.len(), n, "threshold.len() != membrane.len()");
assert_eq!(spikes_out.len(), n, "spikes_out.len() != membrane.len()");
let dt_over_tau = dt_over_tau.clamp(-DT_OVER_TAU_MAX, DT_OVER_TAU_MAX);
for i in 0..n {
let mp = i32::from(membrane[i]);
let leak = i32::from(resting[i]) - mp;
let current_term = (i32::from(input_currents[i]) * i32::from(resistance[i])) / 1000;
let delta_v = (dt_over_tau * (leak + current_term)) / 1000;
let new_v = mp.saturating_add(delta_v).clamp(-100, 50);
membrane[i] = new_v as i16;
spikes_out[i] = new_v >= i32::from(threshold[i]);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn integrate_batch_avx2(
membrane: &mut [i16],
resting: &[i16],
input_currents: &[i16],
resistance: &[i16],
threshold: &[i16],
dt_over_tau: i32,
spikes_out: &mut [bool],
) {
const WIDTH: usize = 16; let dt_over_tau = dt_over_tau.clamp(-DT_OVER_TAU_MAX, DT_OVER_TAU_MAX);
let n = membrane.len();
let chunks = n / WIDTH;
let dt_v = _mm256_set1_epi32(dt_over_tau);
for c in 0..chunks {
let off = c * WIDTH;
let mp = _mm256_loadu_si256(membrane.as_ptr().add(off) as *const __m256i);
let rp = _mm256_loadu_si256(resting.as_ptr().add(off) as *const __m256i);
let ic = _mm256_loadu_si256(input_currents.as_ptr().add(off) as *const __m256i);
let res = _mm256_loadu_si256(resistance.as_ptr().add(off) as *const __m256i);
let th = _mm256_loadu_si256(threshold.as_ptr().add(off) as *const __m256i);
let mp_lo = _mm256_cvtepi16_epi32(_mm256_castsi256_si128(mp));
let mp_hi = _mm256_cvtepi16_epi32(_mm256_extracti128_si256(mp, 1));
let rp_lo = _mm256_cvtepi16_epi32(_mm256_castsi256_si128(rp));
let rp_hi = _mm256_cvtepi16_epi32(_mm256_extracti128_si256(rp, 1));
let ic_lo = _mm256_cvtepi16_epi32(_mm256_castsi256_si128(ic));
let ic_hi = _mm256_cvtepi16_epi32(_mm256_extracti128_si256(ic, 1));
let res_lo = _mm256_cvtepi16_epi32(_mm256_castsi256_si128(res));
let res_hi = _mm256_cvtepi16_epi32(_mm256_extracti128_si256(res, 1));
let new_lo = lif_lane(mp_lo, rp_lo, ic_lo, res_lo, dt_v);
let new_hi = lif_lane(mp_hi, rp_hi, ic_hi, res_hi, dt_v);
let packed = _mm256_packs_epi32(new_lo, new_hi);
let final_mp = _mm256_permute4x64_epi64(packed, 0b1101_1000);
_mm256_storeu_si256(membrane.as_mut_ptr().add(off) as *mut __m256i, final_mp);
let gt = _mm256_cmpgt_epi16(final_mp, th);
let eq = _mm256_cmpeq_epi16(final_mp, th);
let ge = _mm256_or_si256(gt, eq);
let mask = _mm256_movemask_epi8(ge) as u32;
for j in 0..WIDTH {
spikes_out[off + j] = (mask >> (j * 2)) & 1 != 0;
}
}
let tail = chunks * WIDTH;
integrate_batch_scalar(
&mut membrane[tail..],
&resting[tail..],
&input_currents[tail..],
&resistance[tail..],
&threshold[tail..],
dt_over_tau,
&mut spikes_out[tail..],
);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn div1024_toward_zero(x: __m256i) -> __m256i {
let bias = _mm256_and_si256(_mm256_srai_epi32(x, 31), _mm256_set1_epi32(1023));
_mm256_srai_epi32(_mm256_add_epi32(x, bias), 10)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn lif_lane(
mp: __m256i,
rp: __m256i,
ic: __m256i,
res: __m256i,
dt_over_tau: __m256i,
) -> __m256i {
let leak = _mm256_sub_epi32(rp, mp);
let current_scaled = _mm256_mullo_epi32(ic, res);
let current_term = div1024_toward_zero(current_scaled);
let sum = _mm256_add_epi32(leak, current_term);
let delta = _mm256_mullo_epi32(sum, dt_over_tau);
let delta_scaled = div1024_toward_zero(delta);
let new_mp = _mm256_add_epi32(mp, delta_scaled);
_mm256_max_epi32(
_mm256_set1_epi32(-100),
_mm256_min_epi32(_mm256_set1_epi32(50), new_mp),
)
}
#[cfg(all(test, feature = "simd"))]
mod tests {
#![allow(clippy::shadow_unrelated)]
#![allow(clippy::cast_precision_loss)] use super::*;
use crate::lif_neuron::{LIFNeuron, VoltageResolution};
use proptest::prelude::*;
const EQUIV_CURRENT_PRODUCT_MAX: i32 = 100_000;
const EQUIV_DT_OVER_TAU_MAX: i32 = 200;
const EQUIV_TOLERANCE_MV: i32 = 2;
const ARMS: [(i16, i16, i16, i32); 11] = [
(-70, 0, 100, 50),
(-70, 200, 100, 50),
(-70, -200, 100, 50),
(-70, 1000, 100, 50),
(-70, -1000, 100, 50),
(-70, 500, 100, 50),
(-70, -500, 100, 50),
(0, -10000, 10, 50),
(50, 7500, 10, 200),
(37, 1000, 100, 5), (50, 870, 100, 5), ];
type SoaBatch = (Vec<i16>, Vec<i16>, Vec<i16>, Vec<i16>, Vec<i16>);
fn soa_in_domain(raw: &[(i16, i16, i16, i16, i16)]) -> SoaBatch {
let mut membrane = Vec::with_capacity(raw.len());
let mut resting = Vec::with_capacity(raw.len());
let mut current = Vec::with_capacity(raw.len());
let mut resistance = Vec::with_capacity(raw.len());
let mut threshold = Vec::with_capacity(raw.len());
for &(mp, rp, ic_raw, res, th) in raw {
let lim = EQUIV_CURRENT_PRODUCT_MAX / i32::from(res).abs().max(1);
let lim = lim.min(i32::from(i16::MAX)) as i16;
membrane.push(mp);
resting.push(rp);
current.push(ic_raw.clamp(-lim, lim));
resistance.push(res);
threshold.push(th);
}
(membrane, resting, current, resistance, threshold)
}
#[test]
#[cfg(target_arch = "x86_64")]
fn ci_runner_actually_has_avx2() {
if std::env::var_os("NEURALOS_REQUIRE_AVX2").is_some() {
assert!(
matches!(detect_simd_support(), SimdSupport::Avx2),
"CI asked for the AVX2 gate and this runner cannot run it; \
every divergence number in this module went unchecked"
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn simd_membrane_stays_bounded() {
let n = 256;
let mut mp = vec![60i16; n]; let rp = vec![-70i16; n];
let ic = vec![1000i16; n];
let res = vec![100i16; n];
let th = vec![-55i16; n];
let mut spikes = vec![false; n];
let dtot = dt_over_tau(1000, 20_000);
integrate_lif_batch(&mut mp, &rp, &ic, &res, &th, dtot, &mut spikes);
for &v in &mp {
assert!((-100..=50).contains(&v), "out of bounds: {v}");
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn simd_approximates_scalar_within_tolerance() {
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — skipping equivalence test)");
return;
}
let n = 512;
let mut mp_a = vec![0i16; n];
let mut mp_b = vec![0i16; n];
let mut rp = vec![0i16; n];
let mut ic = vec![0i16; n];
let res = vec![100i16; n];
let mut th = vec![0i16; n];
for i in 0..n {
let v = ((i as i32 * 7) % 201 - 100) as i16; mp_a[i] = v;
mp_b[i] = v;
rp[i] = ((i as i32 * 3) % 201 - 100) as i16;
ic[i] = ((i as i32 * 11) % 2001 - 1000) as i16; th[i] = -55 - (i as i16 % 20);
}
let dtot = dt_over_tau(1000, 20_000);
let mut spikes_a = vec![false; n];
let mut spikes_b = vec![false; n];
integrate_batch_scalar(&mut mp_a, &rp, &ic, &res, &th, dtot, &mut spikes_a);
unsafe {
integrate_batch_avx2(&mut mp_b, &rp, &ic, &res, &th, dtot, &mut spikes_b);
}
let mut max_diff = 0i32;
let mut disagree = 0usize;
for i in 0..n {
let d = (i32::from(mp_a[i]) - i32::from(mp_b[i])).abs();
if d > max_diff {
max_diff = d;
}
if spikes_a[i] != spikes_b[i] {
disagree += 1;
}
}
assert!(
max_diff <= 2,
"SIMD diverged from scalar by {max_diff} mV (>2)"
);
let disagree_ratio = disagree as f64 / n as f64;
assert!(
disagree_ratio < 0.10,
"{disagree}/{n} spike disagreements (>10%)"
);
}
#[test]
fn dt_over_tau_max_is_the_documented_bound() {
let max_leak = i64::from(i16::MAX) - i64::from(i16::MIN);
let max_product = i64::from(i16::MIN) * i64::from(i16::MIN);
let max_current_term = max_product / 1000;
let max_sum = max_leak + max_current_term;
assert_eq!(max_leak, 65_535);
assert_eq!(max_product, 1_073_741_824);
assert_eq!(max_sum, 1_139_276);
let bound = i64::from(DT_OVER_TAU_MAX);
assert!(
bound * max_sum <= i64::from(i32::MAX),
"DT_OVER_TAU_MAX is too large: {DT_OVER_TAU_MAX} * {max_sum} overflows i32"
);
let next = DT_OVER_TAU_MAX + 1;
assert!(
(bound + 1) * max_sum > i64::from(i32::MAX),
"DT_OVER_TAU_MAX is not tight: {next} would also fit"
);
}
#[test]
fn dt_over_tau_is_non_negative_and_saturated_over_the_whole_u32_domain() {
assert_eq!(dt_over_tau(2_147_484, u32::MAX), 0);
for &dt in &[0u32, 1, 1000, 10_000, i32::MAX as u32, 2_147_484, u32::MAX] {
for &tau in &[1u32, 20_000, i32::MAX as u32, 2_147_483_648, u32::MAX] {
let v = dt_over_tau(dt, tau);
assert!(
(0..=DT_OVER_TAU_MAX).contains(&v),
"dt_over_tau({dt}, {tau}) = {v} is outside 0..={DT_OVER_TAU_MAX}"
);
}
}
assert_eq!(dt_over_tau(0, 20_000), 0, "tau == 0 guard unchanged");
assert_eq!(dt_over_tau(1000, 0), 0, "tau == 0 guard unchanged");
assert_eq!(
dt_over_tau(1000, 20_000),
50,
"the physical default is untouched"
);
}
#[cfg(target_arch = "x86_64")]
fn measure_divergence(
mp: &[i16],
rp: &[i16],
ic: &[i16],
res: &[i16],
th: &[i16],
dtot: i32,
) -> Option<(i32, usize, usize, usize)> {
let n = mp.len();
assert!(n.is_multiple_of(LANES), "fixture must be all vector lanes");
let mut mp_s = mp.to_vec();
let mut sp_s = vec![false; n];
integrate_batch_scalar(&mut mp_s, rp, ic, res, th, dtot, &mut sp_s);
for &v in &mp_s {
assert!((-100..=50).contains(&v), "scalar left the mV grid: {v}");
}
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
return None;
}
let mut mp_v = mp.to_vec();
let mut sp_v = vec![false; n];
unsafe {
integrate_batch_avx2(&mut mp_v, rp, ic, res, th, dtot, &mut sp_v);
}
for &v in &mp_v {
assert!((-100..=50).contains(&v), "AVX2 left the mV grid: {v}");
}
let max_diff = mp_s
.iter()
.zip(&mp_v)
.map(|(a, b)| (i32::from(*a) - i32::from(*b)).abs())
.max()
.expect("non-empty batch");
let membrane_diffs = mp_s.iter().zip(&mp_v).filter(|(a, b)| a != b).count();
let spike_diffs = sp_s.iter().zip(&sp_v).filter(|(a, b)| a != b).count();
let both_saturating = mp_s
.iter()
.zip(&mp_v)
.filter(|(a, b)| a == b && matches!(**a, -100 | 50))
.count();
Some((max_diff, membrane_diffs, spike_diffs, both_saturating))
}
#[test]
#[cfg(target_arch = "x86_64")]
fn overflow_corners_at_max_dt_over_tau() {
const CORNERS: [i16; 5] = [i16::MIN, -1, 0, 1, i16::MAX];
let (mut mp, mut rp, mut ic, mut res, mut th) =
(Vec::new(), Vec::new(), Vec::new(), Vec::new(), Vec::new());
for &a in &CORNERS {
for &b in &CORNERS {
for &c in &CORNERS {
for &d in &CORNERS {
for &e in &CORNERS {
mp.push(a);
rp.push(b);
ic.push(c);
res.push(d);
th.push(e);
}
}
}
}
}
assert_eq!(mp.len(), 5usize.pow(5), "5^5 corner combinations");
for v in [&mut mp, &mut rp, &mut ic, &mut res, &mut th] {
pad_to_lanes(v);
}
for dtot in [DT_OVER_TAU_MAX, -DT_OVER_TAU_MAX] {
let Some((max_diff, membrane_diffs, spike_diffs, both_saturating)) =
measure_divergence(&mp, &rp, &ic, &res, &th, dtot)
else {
eprintln!("(AVX2 not available — corner agreement not checked)");
return;
};
assert_eq!(
(max_diff, membrane_diffs, spike_diffs),
(4, 180, 0),
"corner divergence moved at dt_over_tau = {dtot}"
);
assert_eq!(
both_saturating, 2371,
"the corner fixture's blind fraction moved at dt_over_tau = {dtot}; \
75.6 % of it cannot see the arithmetic and the doc says so"
);
}
}
#[cfg(target_arch = "x86_64")]
type SoaFixture = (Vec<i16>, Vec<i16>, Vec<i16>, Vec<i16>, Vec<i16>);
#[cfg(target_arch = "x86_64")]
fn mv_grid_fixture() -> SoaFixture {
const MEMBRANE: [i16; 6] = [-100, -70, -55, -20, 0, 50];
const RESTING: [i16; 8] = [-100, -70, -55, -37, -20, 0, 37, 50];
const THRESHOLD: [i16; 7] = [-100, -70, -56, -55, -20, 0, 50];
const OFFSET: [i16; 11] = [-770, -400, -210, -80, -20, 0, 20, 80, 210, 400, 770];
const RESISTANCE: i16 = 100;
let (mut mp, mut rp, mut ic, mut res, mut th) =
(Vec::new(), Vec::new(), Vec::new(), Vec::new(), Vec::new());
for &membrane in &MEMBRANE {
for &resting in &RESTING {
let leak = resting - membrane;
for &offset in &OFFSET {
let input = -10 * leak + offset;
for &threshold in &THRESHOLD {
mp.push(membrane);
rp.push(resting);
ic.push(input);
res.push(RESISTANCE);
th.push(threshold);
}
}
}
}
assert_eq!(mp.len(), 3696, "6 × 8 × 11 × 7 rows");
assert!(
mp.len().is_multiple_of(LANES),
"3696 = 231 × 16, no padding"
);
(mp, rp, ic, res, th)
}
#[cfg(target_arch = "x86_64")]
fn avx2_lane_model(
mp: i16,
rp: i16,
ic: i16,
res: i16,
dtot: i32,
floor_current_term: bool,
) -> i32 {
let mp = i32::from(mp);
let leak = i32::from(rp) - mp;
let product = i32::from(ic) * i32::from(res);
let current_term = if floor_current_term {
product >> 10
} else {
product / 1024
};
let delta = ((leak + current_term) * dtot) / 1024;
(mp + delta).clamp(-100, 50)
}
#[test]
#[cfg(target_arch = "x86_64")]
fn the_recorded_fork_re_measured_against_the_committed_choice() {
let (mp, rp, ic, res, th) = mv_grid_fixture();
let n = mp.len();
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!(
"(AVX2 not available — the model cannot be validated, so nothing is reported)"
);
return;
}
for dtot in [DT_OVER_TAU_MAX, -DT_OVER_TAU_MAX] {
let mut mp_v = mp.clone();
let mut sp_v = vec![false; n];
unsafe {
integrate_batch_avx2(&mut mp_v, &rp, &ic, &res, &th, dtot, &mut sp_v);
}
for i in 0..n {
let modelled = avx2_lane_model(mp[i], rp[i], ic[i], res[i], dtot, false);
assert_eq!(
modelled,
i32::from(mp_v[i]),
"the model disagrees with the real AVX2 kernel at row {i}, \
dt_over_tau = {dtot}: modelled {modelled} vs kernel {}",
mp_v[i],
);
}
}
let measure = |floor_current_term: bool| {
let (mut max_diff, mut membrane_diffs, mut spike_diffs) = (0i32, 0usize, 0usize);
for dtot in [DT_OVER_TAU_MAX, -DT_OVER_TAU_MAX] {
let mut mp_s = mp.clone();
let mut sp_s = vec![false; n];
integrate_batch_scalar(&mut mp_s, &rp, &ic, &res, &th, dtot, &mut sp_s);
for i in 0..n {
let v = avx2_lane_model(mp[i], rp[i], ic[i], res[i], dtot, floor_current_term);
let s = i32::from(mp_s[i]);
max_diff = max_diff.max((s - v).abs());
membrane_diffs += usize::from(s != v);
spike_diffs += usize::from(sp_s[i] != (v >= i32::from(th[i])));
}
}
(max_diff, membrane_diffs, spike_diffs)
};
assert_eq!(
measure(false),
(15, 4592, 157),
"the committed (truncate, truncate) moved"
);
assert_eq!(
measure(true),
(15, 4277, 122),
"the recorded fork (floor current_term, truncate delta) moved"
);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn mv_grid_divergence_at_max_dt_over_tau() {
let (mp, rp, ic, res, th) = mv_grid_fixture();
let expected: [(i32, i32, usize, usize, usize); 2] = [
(DT_OVER_TAU_MAX, 8, 2331, 65, 1225),
(-DT_OVER_TAU_MAX, 15, 2261, 92, 1379),
];
for (dtot, max_diff, membrane_diffs, spike_diffs, both_saturating) in expected {
let Some(got) = measure_divergence(&mp, &rp, &ic, &res, &th, dtot) else {
eprintln!("(AVX2 not available — mV-grid divergence not checked)");
return;
};
assert_eq!(
got,
(max_diff, membrane_diffs, spike_diffs, both_saturating),
"mV-grid divergence moved at dt_over_tau = {dtot}"
);
assert!(
both_saturating * 2 < mp.len(),
"non-saturating rows must dominate: {both_saturating}/{} saturate at {dtot}",
mp.len()
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn a_spike_bit_diverges_on_the_mv_grid_at_max_dt_over_tau() {
let (mp, rp, ic, res, th) = (-100i16, -100i16, 247i16, 100i16, -55i16);
let mut mp_s = vec![mp; LANES];
let mut sp_s = vec![false; LANES];
integrate_batch_scalar(
&mut mp_s,
&[rp; LANES],
&[ic; LANES],
&[res; LANES],
&[th; LANES],
DT_OVER_TAU_MAX,
&mut sp_s,
);
assert_eq!(mp_s[0], -55, "scalar lands on threshold exactly");
assert!(sp_s[0], "scalar fires");
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — spike divergence not checked)");
return;
}
let mut mp_v = vec![mp; LANES];
let mut sp_v = vec![false; LANES];
unsafe {
integrate_batch_avx2(
&mut mp_v,
&[rp; LANES],
&[ic; LANES],
&[res; LANES],
&[th; LANES],
DT_OVER_TAU_MAX,
&mut sp_v,
);
}
assert_eq!(mp_v[0], -56, "AVX2 lands one mV short");
assert!(!sp_v[0], "AVX2 does not fire — the spike bits disagree");
}
#[test]
fn unequal_slice_lengths_panic_before_touching_the_caller() {
const N: usize = 32;
const SHORT: usize = 17; const SENTINEL: i16 = -70;
const NAMES: [&str; 5] = [
"resting",
"input_currents",
"resistance",
"threshold",
"spikes_out",
];
let run = |scalar_entry: bool, lens: [usize; 5]| -> (bool, Vec<i16>, Vec<bool>) {
let mut membrane = vec![SENTINEL; N];
let mut spikes = vec![true; lens[4]];
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let entry = if scalar_entry {
integrate_batch_scalar
} else {
integrate_lif_batch
};
entry(
&mut membrane,
&vec![0i16; lens[0]],
&vec![0i16; lens[1]],
&vec![100i16; lens[2]],
&vec![-55i16; lens[3]],
50,
&mut spikes,
);
}))
.is_err();
(panicked, membrane, spikes)
};
let previous = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let mut failures: Vec<String> = Vec::new();
for (i, name) in NAMES.iter().enumerate() {
for (entry_name, scalar_entry) in [
("integrate_lif_batch", false),
("integrate_batch_scalar", true),
] {
for (label, len) in [("short", SHORT), ("long", N * 2)] {
let mut lens = [N; 5];
lens[i] = len;
let (panicked, membrane, spikes) = run(scalar_entry, lens);
let name = &format!("{name} via {entry_name}");
if !panicked {
let why = if label == "short" {
"no panic"
} else {
"no panic — a prefix was integrated silently"
};
failures.push(format!("{label} `{name}`: {why}"));
}
if let Some(pos) = membrane.iter().position(|&v| v != SENTINEL) {
failures.push(format!(
"{label} `{name}`: membrane[{pos}] = {} was written before the panic",
membrane[pos]
));
}
if let Some(pos) = spikes.iter().position(|&b| !b) {
failures.push(format!(
"{label} `{name}`: spikes_out[{pos}] was written before the panic"
));
}
}
}
}
std::panic::set_hook(previous);
assert!(
failures.is_empty(),
"length contract unenforced:
{}",
failures.join(
"
"
)
);
}
#[test]
fn scalar_matches_itself() {
let n = 64;
let mut a = vec![-70i16; n];
let mut b = vec![-70i16; n];
let rp = vec![-70i16; n];
let ic = vec![200i16; n];
let res = vec![100i16; n];
let th = vec![-55i16; n];
let dtot = dt_over_tau(1000, 20_000);
let mut sa = vec![false; n];
let mut sb = vec![false; n];
integrate_batch_scalar(&mut a, &rp, &ic, &res, &th, dtot, &mut sa);
integrate_batch_scalar(&mut b, &rp, &ic, &res, &th, dtot, &mut sb);
assert_eq!(a, b);
assert_eq!(sa, sb);
}
#[test]
fn scalar_batch_matches_integrate_and_fire_in_the_unclamped_regime() {
let mut compared = 0usize;
let mut unclamped = 0usize;
for &mp in &[-90i16, -70, -55, 0, 40] {
for &rp in &[-100i16, -70, 0, 50] {
for &resistance in &[1i16, 10, 100, 500, 1000] {
for input in (-1500i16..=1500).step_by(37) {
for &(dt_us, tau_us) in
&[(1000u32, 20_000u32), (500, 20_000), (100, 10_000)]
{
let mut n = LIFNeuron::new(0);
n.voltage_resolution = VoltageResolution::Millivolt;
n.membrane_potential = mp;
n.resting_potential = rp;
n.threshold = i16::MAX; n.tau_membrane_us = tau_us;
n.resistance_mohm = resistance as u16;
n.noise_amplitude_ua = 0;
n.synaptic_current_ua = 0;
n.adaptation_current_ua = 0;
n.refractory_time_us = 0;
let fired = n.integrate_and_fire(input, dt_us, 0);
assert!(!fired, "i16::MAX threshold must be unreachable");
let dtot = dt_over_tau(dt_us, tau_us);
let mut membrane = vec![mp];
let mut spikes = vec![false];
integrate_batch_scalar(
&mut membrane,
&[rp],
&[input],
&[resistance],
&[i16::MAX],
dtot,
&mut spikes,
);
assert_eq!(
membrane[0], n.membrane_potential,
"mp={mp} rp={rp} input={input} resistance={resistance} \
dt_us={dt_us} tau_us={tau_us} dt_over_tau={dtot}"
);
assert_eq!(
spikes[0], fired,
"spike bit differs at an unreachable threshold"
);
compared += 1;
if membrane[0] != -100 && membrane[0] != 50 {
unclamped += 1;
}
}
}
}
}
}
assert!(compared >= 5000, "sweep shrank to {compared} rows");
assert!(
unclamped * 4 >= compared,
"only {unclamped}/{compared} rows landed inside the clamp — the sweep has gone blind"
);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn avx2_tail_writes_every_remainder_element() {
const DTOT: i32 = 50; if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — skipping tail test)");
return;
}
for n in [0usize, 1, 15, 16, 17, 31, 32, 33, 47, 48, 257] {
let membrane = vec![-70i16; n];
let resting = vec![0i16; n];
let current = vec![500i16; n];
let resistance = vec![100i16; n];
let threshold = vec![-100i16; n];
let mut mp_v = membrane.clone();
let mut sp_v = vec![false; n];
unsafe {
integrate_batch_avx2(
&mut mp_v,
&resting,
¤t,
&resistance,
&threshold,
DTOT,
&mut sp_v,
);
}
let vector_lanes = (n / 16) * 16;
let expected: Vec<i16> = (0..n)
.map(|i| if i < vector_lanes { -65 } else { -64 })
.collect();
assert_eq!(
mp_v,
expected,
"n={n}: {vector_lanes} vector lanes then a {} element tail",
n - vector_lanes
);
assert!(
sp_v.iter().all(|&b| b),
"n={n}: a spike bit was never written — got {sp_v:?}"
);
if n > 0 {
assert_ne!(mp_v[n - 1], -70, "n={n}: the LAST element was not written");
assert!(sp_v[n - 1], "n={n}: the LAST spike bit was not written");
}
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn avx2_and_scalar_trajectories_stay_bounded() {
const N: usize = 16;
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — skipping trajectory test)");
return;
}
let run = |resting: i16, drive: i16, resistance: i16, dtot: i32, steps: usize| {
let rp = vec![resting; N];
let ic = vec![drive; N];
let res = vec![resistance; N];
let th = vec![i16::MAX; N]; let mut mp_s = vec![-70i16; N];
let mut mp_v = vec![-70i16; N];
let mut sp_s = vec![false; N];
let mut sp_v = vec![false; N];
let mut worst_step = 0i32;
for _ in 0..steps {
integrate_batch_scalar(&mut mp_s, &rp, &ic, &res, &th, dtot, &mut sp_s);
integrate_lif_batch(&mut mp_v, &rp, &ic, &res, &th, dtot, &mut sp_v);
let d = (i32::from(mp_s[0]) - i32::from(mp_v[0])).abs();
if d > worst_step {
worst_step = d;
}
}
((i32::from(mp_s[0]) - i32::from(mp_v[0])).abs(), worst_step)
};
let mut worst_end = 0i32;
let mut worst_traj = 0i32;
for &(resting, drive, resistance, dtot) in &ARMS {
let (short_end, short_worst) = run(resting, drive, resistance, dtot, 1_000);
let (long_end, long_worst) = run(resting, drive, resistance, dtot, 20_000);
assert_eq!(
(short_end, short_worst),
(long_end, long_worst),
"arm rp={resting} drive={drive} res={resistance} dt_over_tau={dtot}: the gap \
grew between 1_000 and 20_000 steps, so it is accumulating, not parking"
);
worst_end = worst_end.max(long_end);
worst_traj = worst_traj.max(long_worst);
}
let named: Vec<i32> = [0i16, 200, -200]
.iter()
.map(|&drive| run(-70, drive, 100, 50, 10_000).0)
.collect();
assert_eq!(
named,
vec![0, 1, 1],
"the named arms moved (they were 0, 1, 19 before the fix)"
);
assert_eq!(
(worst_end, worst_traj),
(8, 8),
"trajectory bounds moved; the module doc records 8 mV domain-wide, at \
the fixed point and at any step, established by exhaustive sweep"
);
}
const LANES: usize = 16;
fn pad_to_lanes<T: Copy>(v: &mut Vec<T>) {
let last = *v.last().expect("fixture must be non-empty");
while !v.len().is_multiple_of(LANES) {
v.push(last);
}
}
fn current_term_class_representatives() -> Vec<(i16, i16)> {
fn factor_i16(p: i32) -> Option<(i16, i16)> {
if p.abs() <= i32::from(i16::MAX) {
return Some((p as i16, 1));
}
(2..=64i32).find_map(|d| {
(p % d == 0 && (p / d).abs() <= i32::from(i16::MAX))
.then(|| ((p / d) as i16, d as i16))
})
}
let mut seen = std::collections::BTreeSet::new();
let mut out = Vec::new();
for p in -EQUIV_CURRENT_PRODUCT_MAX..=EQUIV_CURRENT_PRODUCT_MAX {
let class = (p / 1000, p / 1024);
if !seen.insert(class) {
continue;
}
if let Some(pair) = factor_i16(p) {
out.push(pair);
} else {
seen.remove(&class); }
}
out
}
#[cfg(target_arch = "x86_64")]
fn max_membrane_gap(
start: &[i16],
resting: &[i16],
current: &[i16],
resistance: &[i16],
threshold: &[i16],
dtot: i32,
scratch: &mut (Vec<i16>, Vec<i16>, Vec<bool>, Vec<bool>),
) -> i32 {
let (mp_s, mp_v, sp_s, sp_v) = scratch;
mp_s.copy_from_slice(start);
mp_v.copy_from_slice(start);
integrate_batch_scalar(mp_s, resting, current, resistance, threshold, dtot, sp_s);
unsafe {
integrate_batch_avx2(mp_v, resting, current, resistance, threshold, dtot, sp_v);
}
mp_s.iter()
.zip(mp_v.iter())
.map(|(a, b)| (i32::from(*a) - i32::from(*b)).abs())
.max()
.unwrap_or(0)
}
#[test]
#[ignore = "exhaustive sweep, minutes not milliseconds"]
#[cfg(target_arch = "x86_64")]
fn sweep_reproduces_the_documented_equivalence_domain() {
const DT_SCAN: i32 = 260; const W: usize = 16;
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — skipping the equivalence sweep)");
return;
}
let reps = current_term_class_representatives();
assert_eq!(
reps.len(),
395,
"current-term class coverage changed; the doc says 395 classes"
);
let mut start: Vec<i16> = Vec::new();
let mut resting: Vec<i16> = Vec::new();
for m in -100..=50i16 {
for r in -100..=50i16 {
start.push(m);
resting.push(r);
}
}
pad_to_lanes(&mut start);
pad_to_lanes(&mut resting);
let n = start.len();
assert!(n.is_multiple_of(LANES), "fixture must be all vector lanes");
let threshold = vec![i16::MAX; n];
let mut scratch = (vec![0i16; n], vec![0i16; n], vec![false; n], vec![false; n]);
let mut max_by_dt = vec![0i32; (DT_SCAN + 1) as usize];
for &(ic0, res0) in &reps {
let current = vec![ic0; n];
let resistance = vec![res0; n];
for dt in 0..=DT_SCAN {
let g = max_membrane_gap(
&start,
&resting,
¤t,
&resistance,
&threshold,
dt,
&mut scratch,
);
let slot = &mut max_by_dt[dt as usize];
if g > *slot {
*slot = g;
}
}
}
let in_domain = max_by_dt[..=(EQUIV_DT_OVER_TAU_MAX as usize)]
.iter()
.copied()
.max()
.expect("non-empty");
assert_eq!(
in_domain, EQUIV_TOLERANCE_MV,
"the equivalence domain's maximum is not the documented {EQUIV_TOLERANCE_MV} mV"
);
assert_eq!(max_by_dt[1], 0, "doc says dt_over_tau = 1 gives 0");
assert_eq!(max_by_dt[50], 1, "doc says dt_over_tau = 50 gives 1");
assert_eq!(
max_by_dt[200], 2,
"doc says the bound is tight at dt_over_tau = 200"
);
let first_three = max_by_dt
.iter()
.position(|&g| g >= 3)
.expect("the scan must reach 3 before dt_over_tau = 260");
assert_eq!(
first_three, 228,
"the doc names 228 as the first dt_over_tau reaching 3"
);
let mut wide = (vec![0i16; W], vec![0i16; W], vec![false; W], vec![false; W]);
let gap = max_membrane_gap(
&[-100; W],
&[50; W],
&[1000; W],
&[100; W],
&[i16::MAX; W],
228,
&mut wide,
);
assert_eq!(
gap, 3,
"the named witness (membrane -100, resting 50, P = 100_000) must give 3"
);
}
#[cfg(target_arch = "x86_64")]
fn arms_worst_gaps() -> (i32, i32) {
ARMS.iter()
.map(|&(rp, drive, res, dt)| {
let rp_v = vec![rp; LANES];
let ic_v = vec![drive; LANES];
let res_v = vec![res; LANES];
let th_v = vec![i16::MAX; LANES];
let mut a = vec![-70i16; LANES];
let mut b = vec![-70i16; LANES];
let (mut sa, mut sb) = (vec![false; LANES], vec![false; LANES]);
let mut traj = 0i32;
for _ in 0..20_000 {
integrate_batch_scalar(&mut a, &rp_v, &ic_v, &res_v, &th_v, dt, &mut sa);
integrate_lif_batch(&mut b, &rp_v, &ic_v, &res_v, &th_v, dt, &mut sb);
let d = (i32::from(a[0]) - i32::from(b[0])).abs();
if d > traj {
traj = d;
}
}
((i32::from(a[0]) - i32::from(b[0])).abs(), traj)
})
.fold((0i32, 0i32), |acc, x| (acc.0.max(x.0), acc.1.max(x.1)))
}
#[test]
#[ignore = "exhaustive sweep, minutes not milliseconds"]
#[cfg(target_arch = "x86_64")]
fn sweep_reproduces_the_documented_trajectory_maxima() {
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
eprintln!("(AVX2 not available — skipping the trajectory sweep)");
return;
}
let reps = current_term_class_representatives();
assert_eq!(reps.len(), 395, "current-term class coverage changed");
let mut resting: Vec<i16> = (-100..=50i16).collect();
pad_to_lanes(&mut resting);
let n = resting.len();
assert!(n.is_multiple_of(LANES), "fixture must be all vector lanes");
let threshold = vec![i16::MAX; n];
let (mut mp_s, mut mp_v) = (vec![0i16; n], vec![0i16; n]);
let (mut sp_s, mut sp_v) = (vec![false; n], vec![false; n]);
let mut worst_end = 0i32;
let mut worst_traj = 0i32;
let mut extremal = (0i16, 0i16, 0i16, 0i32);
for &(ic0, res0) in &reps {
let current = vec![ic0; n];
let resistance = vec![res0; n];
for dt in 0..=EQUIV_DT_OVER_TAU_MAX {
mp_s.fill(-70);
mp_v.fill(-70);
let mut local_traj = vec![0i32; n];
let mut settled = 0;
let mut prev_s = vec![0i16; n];
let mut prev_v = vec![0i16; n];
for _ in 0..400 {
prev_s.copy_from_slice(&mp_s);
prev_v.copy_from_slice(&mp_v);
integrate_batch_scalar(
&mut mp_s,
&resting,
¤t,
&resistance,
&threshold,
dt,
&mut sp_s,
);
unsafe {
integrate_batch_avx2(
&mut mp_v,
&resting,
¤t,
&resistance,
&threshold,
dt,
&mut sp_v,
);
}
for i in 0..n {
let d = (i32::from(mp_s[i]) - i32::from(mp_v[i])).abs();
if d > local_traj[i] {
local_traj[i] = d;
}
}
if prev_s == mp_s && prev_v == mp_v {
settled += 1;
if settled == 2 {
break;
}
} else {
settled = 0;
}
}
for i in 0..n {
let end = (i32::from(mp_s[i]) - i32::from(mp_v[i])).abs();
if end > worst_end {
worst_end = end;
extremal = (resting[i], ic0, res0, dt);
}
if local_traj[i] > worst_traj {
worst_traj = local_traj[i];
}
}
}
}
assert_eq!(
(worst_end, worst_traj),
(8, 8),
"the doc records 8 mV domain-wide, at the fixed point and at any step"
);
let (arms_end, arms_traj) = arms_worst_gaps();
assert_eq!(
(arms_end, arms_traj),
(worst_end, worst_traj),
"ARMS tops out at ({arms_end}, {arms_traj}) while the domain reaches \
({worst_end}, {worst_traj}), so the fast test is pinning its own arm set rather \
than the domain (the sweep's own extremal was resting {}, P = {}, dt_over_tau {})",
extremal.0,
i32::from(extremal.1) * i32::from(extremal.2),
extremal.3
);
}
#[test]
fn the_batch_diverges_from_the_neuron_by_exactly_its_own_clamp() {
const CASES: [(u32, u32, i16, i16, i16, i16); 5] = [
(40_000, 20_000, -100, -70, 0, 100), (1000, 500, -70, -70, 200, 100), (1000, 400, -70, -70, -100, 100), (20_000, 1000, -70, -70, 10, 100), (1000, 531, -70, -70, 200, 100), ];
for (dt_us, tau_us, mp, rp, input, resistance) in CASES {
let exact = crate::lif_neuron::dt_over_tau(dt_us, tau_us);
let clamped = dt_over_tau(dt_us, tau_us);
let neuron = |tau: u32, dt: u32| {
let mut n = LIFNeuron::new(0);
n.voltage_resolution = VoltageResolution::Millivolt;
n.membrane_potential = mp;
n.resting_potential = rp;
n.threshold = i16::MAX; n.tau_membrane_us = tau;
n.resistance_mohm = resistance as u16;
n.noise_amplitude_ua = 0;
n.synaptic_current_ua = 0;
n.adaptation_current_ua = 0;
n.refractory_time_us = 0;
let _ = n.integrate_and_fire(input, dt, 0);
n.membrane_potential
};
let mut membrane = vec![mp];
let mut spikes = vec![false];
integrate_batch_scalar(
&mut membrane,
&[rp],
&[input],
&[resistance],
&[i16::MAX],
clamped,
&mut spikes,
);
let at_clamped = neuron(1_000_000, 1_884_000);
let expected = if exact <= i64::from(DT_OVER_TAU_MAX) {
neuron(tau_us, dt_us)
} else {
at_clamped
};
assert_eq!(
membrane[0], expected,
"dt={dt_us} tau={tau_us}: batch {} vs expected {expected} \
(exact dt_over_tau {exact}, clamped {clamped})",
membrane[0],
);
if exact <= i64::from(DT_OVER_TAU_MAX) {
assert_eq!(
membrane[0],
neuron(tau_us, dt_us),
"inside the clamp the two must be bit-equal"
);
}
}
let mut n = LIFNeuron::new(0);
n.voltage_resolution = VoltageResolution::Millivolt;
n.membrane_potential = -100;
n.resting_potential = -70;
n.threshold = i16::MAX;
n.tau_membrane_us = 20_000;
n.noise_amplitude_ua = 0;
let _ = n.integrate_and_fire(0, 40_000, 0);
assert_eq!(n.membrane_potential, -40, "neuron: 2000 * 30 / 1000 = +60");
let mut membrane = vec![-100i16];
let mut spikes = vec![false];
integrate_batch_scalar(
&mut membrane,
&[-70],
&[0],
&[100],
&[i16::MAX],
dt_over_tau(40_000, 20_000),
&mut spikes,
);
assert_eq!(membrane[0], -44, "batch: 1884 * 30 / 1000 = +56");
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(512))]
#[test]
fn prop_scalar_batch_is_bit_equal_to_integrate_and_fire(
mp in -100i16..=50,
rp in -100i16..=50,
th in -100i16..=50,
(input, resistance) in prop_oneof![
(any::<i16>(), 0i16..=i16::MAX),
(-2000i16..=2000, 0i16..=1000),
],
tau_us in 1u32..=1_000_000,
dt_ratio in 0u32..=10_000,
) {
let dt_us = u32::try_from(u64::from(tau_us) * u64::from(dt_ratio) / 1000)
.expect("tau_us <= 1e6 and dt_ratio <= 1e4, so dt_us <= 1e7");
let fixture = |threshold: i16| {
let mut n = LIFNeuron::new(0);
n.voltage_resolution = VoltageResolution::Millivolt;
n.membrane_potential = mp;
n.resting_potential = rp;
n.threshold = threshold;
n.tau_membrane_us = tau_us;
n.resistance_mohm = resistance as u16;
n.noise_amplitude_ua = 0;
n.synaptic_current_ua = 0;
n.adaptation_current_ua = 0;
n.refractory_time_us = 0;
n
};
let exact = crate::lif_neuron::dt_over_tau(dt_us, tau_us);
let dtot = dt_over_tau(dt_us, tau_us);
prop_assert_eq!(
i64::from(dtot), exact.min(i64::from(DT_OVER_TAU_MAX)),
"simd::dt_over_tau is not lif_neuron::dt_over_tau plus this kernel's clamp"
);
let inside_the_clamp = exact <= i64::from(DT_OVER_TAU_MAX);
let mut quiet = fixture(i16::MAX);
let fired_quiet = quiet.integrate_and_fire(input, dt_us, 0);
prop_assert!(!fired_quiet, "i16::MAX threshold must be unreachable");
let mut live = fixture(th);
let fired = live.integrate_and_fire(input, dt_us, 0);
let mut membrane = vec![mp];
let mut spikes = vec![false];
integrate_batch_scalar(
&mut membrane, &[rp], &[input], &[resistance], &[th], dtot, &mut spikes,
);
if inside_the_clamp {
prop_assert_eq!(
membrane[0], quiet.membrane_potential,
"membrane differs: batch {} vs integrate_and_fire {} \
(mp={} rp={} input={} resistance={} dt_over_tau={})",
membrane[0], quiet.membrane_potential, mp, rp, input, resistance, dtot
);
prop_assert_eq!(
spikes[0], fired,
"spike bit differs at threshold {} with membrane {}",
th, quiet.membrane_potential
);
} else {
let mut clamped = LIFNeuron::new(0);
clamped.voltage_resolution = VoltageResolution::Millivolt;
clamped.membrane_potential = mp;
clamped.resting_potential = rp;
clamped.threshold = i16::MAX;
clamped.resistance_mohm = resistance as u16;
clamped.noise_amplitude_ua = 0;
clamped.synaptic_current_ua = 0;
clamped.adaptation_current_ua = 0;
clamped.refractory_time_us = 0;
clamped.tau_membrane_us = 1_000_000;
let _ = clamped.integrate_and_fire(input, 1_884_000, 0);
prop_assert_eq!(
membrane[0], clamped.membrane_potential,
"above the clamp the batch must equal the neuron run at \
dt_over_tau = DT_OVER_TAU_MAX, not something else"
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn prop_avx2_matches_scalar_in_the_equivalence_domain(
raw in prop::collection::vec(
(-100i16..=50, -100i16..=50, any::<i16>(), any::<i16>(), -100i16..=50),
0..=257,
),
dtot in 0i32..=EQUIV_DT_OVER_TAU_MAX,
) {
if !matches!(detect_simd_support(), SimdSupport::Avx2) {
return Ok(());
}
let (membrane, resting, current, resistance, threshold) = soa_in_domain(&raw);
let n = membrane.len();
let mut mp_s = membrane.clone();
let mut sp_s = vec![false; n];
integrate_batch_scalar(
&mut mp_s, &resting, ¤t, &resistance, &threshold, dtot, &mut sp_s,
);
let mut mp_v = membrane.clone();
let mut sp_v = vec![false; n];
unsafe {
integrate_batch_avx2(
&mut mp_v, &resting, ¤t, &resistance, &threshold, dtot, &mut sp_v,
);
}
for i in 0..n {
let diff = (i32::from(mp_s[i]) - i32::from(mp_v[i])).abs();
prop_assert!(
diff <= EQUIV_TOLERANCE_MV,
"neuron {i}: scalar {} vs avx2 {} differ by {diff} mV (> {EQUIV_TOLERANCE_MV}); \
mp={} rp={} ic={} res={} dt_over_tau={dtot}",
mp_s[i], mp_v[i], membrane[i], resting[i], current[i], resistance[i],
);
if sp_s[i] != sp_v[i] {
let margin = (i32::from(mp_s[i]) - i32::from(threshold[i])).abs();
prop_assert!(
margin <= EQUIV_TOLERANCE_MV,
"neuron {i}: spike disagreement {} vs {} with the scalar membrane {} \
a full {margin} mV from threshold {} — outside the edge band",
sp_s[i], sp_v[i], mp_s[i], threshold[i],
);
}
}
}
}
}