use crate::prelude::*;
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct RateEncoder {
base_rate: f32,
max_rate: f32,
range: (f32, f32),
dt_seconds: f32,
#[cfg_attr(feature = "serde", serde(rename = "accumulators"))]
phases: Vec<f64>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Vec::is_empty")
)]
pending_spikes: Vec<u64>,
}
impl RateEncoder {
pub const DEFAULT_DT_SECONDS: f32 = 0.1;
pub fn new(base_rate: f32, max_rate: f32, range: (f32, f32)) -> Self {
Self::try_new(base_rate, max_rate, range, Self::DEFAULT_DT_SECONDS)
.expect("invalid RateEncoder configuration")
}
pub fn try_new(
base_rate: f32,
max_rate: f32,
range: (f32, f32),
dt_seconds: f32,
) -> Result<Self, EncoderError> {
crate::error::validate_non_negative_finite("base_rate", base_rate)?;
crate::error::validate_non_negative_finite("max_rate", max_rate)?;
if base_rate > max_rate {
return Err(EncoderError::RateOrder);
}
crate::error::validate_range_f32_span("range", range)?;
Self::validate_dt_seconds(dt_seconds)?;
Ok(Self {
base_rate,
max_rate,
range,
dt_seconds,
phases: Vec::new(),
pending_spikes: Vec::new(),
})
}
pub fn dt_seconds(&self) -> f32 {
self.dt_seconds
}
pub fn default_dt_seconds() -> f32 {
Self::DEFAULT_DT_SECONDS
}
fn validate_dt_seconds(dt_seconds: f32) -> Result<(), EncoderError> {
if dt_seconds.is_finite() && dt_seconds > 0.0 {
Ok(())
} else {
Err(EncoderError::NonPositiveOrNonFinite {
parameter: "dt_seconds",
})
}
}
fn normalize(&self, value: f32) -> f32 {
((value - self.range.0) / (self.range.1 - self.range.0)).clamp(0.0, 1.0)
}
fn effective_rate_hz(&self, value: f32, rate_scale: f32) -> f32 {
if !value.is_finite() {
return 0.0;
}
let normalized = f64::from(self.normalize(value));
let base = f64::from(self.base_rate)
+ normalized * (f64::from(self.max_rate) - f64::from(self.base_rate));
let rate = base * f64::from(rate_scale);
if !rate.is_finite() {
return if rate > 0.0 { f32::MAX } else { 0.0 };
}
rate.clamp(0.0, f64::from(f32::MAX)) as f32
}
fn ensure_accumulators(&mut self, num_channels: usize) {
if self.phases.len() < num_channels {
self.phases.resize(num_channels, 0.0);
self.pending_spikes.resize(num_channels, 0);
}
}
fn split_whole_and_frac(value: f64) -> (u64, f64) {
debug_assert!(value.is_finite() && value >= 0.0);
if value < 1.0 {
return (0, value);
}
if value >= u64::MAX as f64 {
return (u64::MAX, 0.0);
}
let whole = value.trunc() as u64;
let frac = (value - whole as f64).clamp(0.0, 1.0 - f64::EPSILON);
(whole, frac)
}
fn apply_streaming_increment(&mut self, channel_idx: usize, increment: f64) {
if increment <= 0.0 {
return;
}
let sum = (self.phases[channel_idx] + increment).min(u64::MAX as f64);
let (whole, frac) = Self::split_whole_and_frac(sum);
self.pending_spikes[channel_idx] = self.pending_spikes[channel_idx].saturating_add(whole);
self.phases[channel_idx] = frac;
}
fn encode_with_rate_scale(&mut self, input: &[f32], rate_scale: f32) -> EncodedOutput {
let mut output = EncodedOutput::new();
if input.is_empty() {
return output;
}
if !rate_scale.is_finite() || rate_scale <= 0.0 {
return output;
}
let mut rng = rand::rng();
for (i, &value) in input.iter().enumerate() {
let Ok(channel) = u16::try_from(i) else {
break;
};
let rate = self.effective_rate_hz(value, rate_scale);
let probability = crate::poisson::probability_from_rate_hz(rate, self.dt_seconds);
if crate::rng::gen_unit_f32_with_rng(&mut rng) < probability {
output.spikes.push(SpikeEvent {
channel,
timestamp: 0,
polarity: true,
});
}
}
output
}
const MAX_SPIKES_PER_CHANNEL_PER_STEP: usize = 1024;
fn streaming_increment(&self, value: f32, rate_scale: f32) -> f64 {
let rate_hz = self.effective_rate_hz(value, rate_scale);
if rate_hz <= 0.0 {
return 0.0;
}
(f64::from(rate_hz) * f64::from(self.dt_seconds)).clamp(0.0, u64::MAX as f64)
}
fn emit_capped_channel_spikes(
&mut self,
channel: u16,
channel_idx: usize,
output: &mut EncodedOutput,
) {
let pending = self.pending_spikes[channel_idx];
if pending == 0 {
return;
}
let emit = pending.min(Self::MAX_SPIKES_PER_CHANNEL_PER_STEP as u64) as usize;
for _ in 0..emit {
output.spikes.push(SpikeEvent {
channel,
timestamp: 0,
polarity: true,
});
}
self.pending_spikes[channel_idx] = pending - emit as u64;
}
fn rate_scale_is_active(rate_scale: f32) -> bool {
rate_scale.is_finite() && rate_scale > 0.0
}
fn encode_step_with_rate_scale(&mut self, input: &[f32], rate_scale: f32) -> EncodedOutput {
let mut output = EncodedOutput::new();
if input.is_empty() {
return output;
}
self.ensure_accumulators(input.len());
let active = Self::rate_scale_is_active(rate_scale);
for (i, &value) in input.iter().enumerate() {
let Ok(channel) = u16::try_from(i) else {
break;
};
if !active {
self.pending_spikes[i] = 0;
self.phases[i] = 0.0;
continue;
}
let increment = self.streaming_increment(value, rate_scale);
self.apply_streaming_increment(i, increment);
self.emit_capped_channel_spikes(channel, i, &mut output);
}
output
}
pub fn encode_with_modulators(
&mut self,
input: &[f32],
modulators: &NeuroModulators,
gain_curves: &NeuromodulatorGainCurves,
) -> EncodedOutput {
<Self as ModulatedEncoder>::encode_with_modulators(self, input, modulators, gain_curves)
}
pub fn encode_step_with_modulators(
&mut self,
input: &[f32],
modulators: &NeuroModulators,
gain_curves: &NeuromodulatorGainCurves,
) -> EncodedOutput {
<Self as ModulatedEncoder>::encode_step_with_modulators(
self,
input,
modulators,
gain_curves,
)
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for RateEncoder {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Helper {
base_rate: f32,
max_rate: f32,
range: (f32, f32),
#[serde(default = "RateEncoder::default_dt_seconds")]
dt_seconds: f32,
#[serde(default)]
accumulators: Vec<f64>,
#[serde(default)]
pending_spikes: Vec<u64>,
}
let helper = Helper::deserialize(deserializer)?;
let mut encoder = Self::try_new(
helper.base_rate,
helper.max_rate,
helper.range,
helper.dt_seconds,
)
.map_err(serde::de::Error::custom)?;
if helper
.accumulators
.iter()
.any(|value| !value.is_finite() || *value < 0.0)
{
return Err(serde::de::Error::custom(
"accumulators must be finite and non-negative",
));
}
let n = helper.accumulators.len().max(helper.pending_spikes.len());
encoder.phases = vec![0.0; n];
encoder.pending_spikes = vec![0; n];
for i in 0..n {
let combined = helper.accumulators.get(i).copied().unwrap_or(0.0);
let (whole_from_acc, phase) = Self::split_whole_and_frac(combined);
let pending = helper.pending_spikes.get(i).copied().unwrap_or(0);
encoder.phases[i] = phase;
encoder.pending_spikes[i] = pending.saturating_add(whole_from_acc);
}
Ok(encoder)
}
}
impl Encoder for RateEncoder {
fn encode(&mut self, input: &[f32]) -> EncodedOutput {
self.encode_with_rate_scale(input, 1.0)
}
fn encode_step(&mut self, input: &[f32]) -> EncodedOutput {
self.encode_step_with_rate_scale(input, 1.0)
}
fn reset(&mut self) {
self.phases.fill(0.0);
self.pending_spikes.fill(0);
}
}
impl ModulatedEncoder for RateEncoder {
fn encode_with_gains(&mut self, input: &[f32], gains: EncodingGains) -> EncodedOutput {
self.encode_with_rate_scale(input, gains.sanitize().firing_rate_scale)
}
fn encode_step_with_gains(&mut self, input: &[f32], gains: EncodingGains) -> EncodedOutput {
self.encode_step_with_rate_scale(input, gains.sanitize().firing_rate_scale)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rate_encoder_basic() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 100.0));
let input = [0.0, 50.0, 100.0];
let output = encoder.encode(&input);
assert!(output.spikes.len() <= 3);
}
#[test]
fn test_rate_encoder_encode_step() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
let output = encoder.encode_step(&[1.0]);
assert_eq!(output.spikes.len(), 1);
let output2 = encoder.encode_step(&[0.5]);
assert_eq!(output2.spikes.len(), 0);
let output3 = encoder.encode_step(&[0.5]);
assert_eq!(output3.spikes.len(), 1);
}
#[test]
fn test_rate_encoder_empty_input() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 100.0));
let input: [f32; 0] = [];
let output = encoder.encode(&input);
assert_eq!(output.spikes.len(), 0);
let output_step = encoder.encode_step(&input);
assert_eq!(output_step.spikes.len(), 0);
}
#[test]
fn test_rate_encoder_single_channel() {
let mut encoder = RateEncoder::new(5.0, 10.0, (0.0, 1.0));
let input = [0.5];
let output = encoder.encode(&input);
assert!(output.spikes.len() <= 1);
}
#[test]
fn test_rate_encoder_below_min() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 100.0));
let input = [-50.0, -100.0, -1.0];
let output = encoder.encode(&input);
assert!(
output.spikes.is_empty(),
"Below-min inputs should produce no spikes"
);
}
#[test]
fn test_rate_encoder_above_max() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 100.0));
let input = [150.0, 200.0, 101.0];
let output = encoder.encode(&input);
assert!(output.spikes.len() <= 3);
for spike in &output.spikes {
assert!(u32::from(spike.channel) < 3);
}
}
#[test]
fn test_rate_encoder_reset_does_not_panic() {
let mut encoder = RateEncoder::new(5.0, 10.0, (0.0, 1.0));
let input = [0.5; 10];
encoder.encode(&input);
encoder.reset();
encoder.encode(&input);
}
#[test]
fn test_rate_encoder_never_panics() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 100.0));
let inputs: [&[f32]; 4] = [&[], &[0.0], &[50.0, 100.0], &[f32::MIN, f32::MAX]];
for input in inputs {
let result =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| encoder.encode(input)));
assert!(result.is_ok());
}
}
#[test]
fn test_rate_encoder_modulated_step_scales_firing_rate() {
let mut encoder = RateEncoder::new(0.0, 5.0, (0.0, 1.0));
let modulators = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let gain_curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
firing_rate: Some(GainCurve::new((0.0, 1.0), (1.0, 2.0))),
..Default::default()
},
..Default::default()
};
let baseline = encoder.encode_step(&[1.0]);
assert!(baseline.spikes.is_empty());
encoder.reset();
let boosted = encoder.encode_step_with_modulators(&[1.0], &modulators, &gain_curves);
assert_eq!(boosted.spikes.len(), 1);
}
#[test]
fn test_rate_encoder_encode_with_modulators() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
let modulators = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let gain_curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
firing_rate: Some(GainCurve::new((0.0, 1.0), (1.0, 2.0))),
..Default::default()
},
..Default::default()
};
let boosted = encoder.encode_step_with_modulators(&[1.0], &modulators, &gain_curves);
assert_eq!(boosted.spikes.len(), 2);
assert!(boosted.spikes.iter().all(|s| s.channel == 0));
let mut baseline = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
let identity = baseline.encode_step_with_modulators(
&[1.0],
&NeuroModulators::default(),
&NeuromodulatorGainCurves::default(),
);
assert_eq!(identity.spikes.len(), 1);
}
#[test]
fn test_rate_encoder_step_shorter_input() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
let _ = encoder.encode_step(&[0.0, 0.0]);
let output = encoder.encode_step(&[1.0]);
assert_eq!(output.spikes.len(), 1);
let quiet = encoder.encode_step(&[0.0, 0.0]);
assert!(quiet.spikes.is_empty());
}
#[test]
fn test_rate_encoder_zero_rate_scale_never_accumulates() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
for _ in 0..10_000 {
let output = encoder.encode_step_with_rate_scale(&[1.0], 0.0);
assert!(
output.spikes.is_empty(),
"zero firing-rate scale must fully silence streaming output"
);
}
}
#[test]
fn test_rate_encoder_non_finite_rate_scale_silences() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
for scale in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY, -1.0] {
let batch = encoder.encode_with_rate_scale(&[1.0], scale);
assert!(
batch.spikes.is_empty(),
"non-finite/negative rate_scale ({scale}) must silence batch encode"
);
let step = encoder.encode_step_with_rate_scale(&[1.0], scale);
assert!(
step.spikes.is_empty(),
"non-finite/negative rate_scale ({scale}) must silence streaming encode"
);
}
encoder.reset();
let recovered = encoder.encode_step_with_rate_scale(&[1.0], 1.0);
assert_eq!(recovered.spikes.len(), 1);
}
#[test]
fn test_rate_encoder_try_new_validation() {
let dt = RateEncoder::DEFAULT_DT_SECONDS;
assert_eq!(
RateEncoder::try_new(f32::NAN, 1.0, (0.0, 1.0), dt).err(),
Some(EncoderError::NonNegativeFinite {
parameter: "base_rate"
})
);
assert_eq!(
RateEncoder::try_new(0.0, f32::INFINITY, (0.0, 1.0), dt).err(),
Some(EncoderError::NonNegativeFinite {
parameter: "max_rate"
})
);
assert_eq!(
RateEncoder::try_new(-5.0, 10.0, (0.0, 1.0), dt).err(),
Some(EncoderError::NonNegativeFinite {
parameter: "base_rate"
})
);
assert_eq!(
RateEncoder::try_new(2.0, 1.0, (0.0, 1.0), dt).err(),
Some(EncoderError::RateOrder)
);
assert_eq!(
RateEncoder::try_new(0.0, 1.0, (1.0, 1.0), dt).err(),
Some(EncoderError::InvalidRange { parameter: "range" })
);
assert_eq!(
RateEncoder::try_new(0.0, 1.0, (f32::MIN, f32::MAX), dt).err(),
Some(EncoderError::InvalidRange { parameter: "range" })
);
assert_eq!(
RateEncoder::try_new(0.0, 10.0, (0.0, 1.0), 0.0).err(),
Some(EncoderError::NonPositiveOrNonFinite {
parameter: "dt_seconds"
})
);
}
#[test]
fn test_rate_encoder_try_new_validates_dt_seconds() {
assert!(RateEncoder::try_new(0.0, 10.0, (0.0, 1.0), 0.001).is_ok());
for dt in [0.0, -0.001, f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
assert!(
RateEncoder::try_new(0.0, 10.0, (0.0, 1.0), dt).is_err(),
"dt_seconds={dt:?} should be rejected"
);
}
}
#[cfg(feature = "serde")]
#[test]
fn test_rate_encoder_serde_rejects_out_of_range_accumulators() {
let backlog = r#"{"base_rate":0.0,"max_rate":10.0,"range":[0.0,1.0],"accumulators":[5.0]}"#;
let res: Result<RateEncoder, _> = serde_json::from_str(backlog);
assert!(res.is_ok());
let negative =
r#"{"base_rate":0.0,"max_rate":10.0,"range":[0.0,1.0],"accumulators":[-0.1]}"#;
let res: Result<RateEncoder, _> = serde_json::from_str(negative);
assert!(res.is_err());
let non_finite =
r#"{"base_rate":0.0,"max_rate":10.0,"range":[0.0,1.0],"accumulators":[null]}"#;
let res: Result<RateEncoder, _> = serde_json::from_str(non_finite);
assert!(res.is_err());
let ok = r#"{"base_rate":0.0,"max_rate":10.0,"range":[0.0,1.0],"dt_seconds":0.1,"accumulators":[0.5]}"#;
let res: Result<RateEncoder, _> = serde_json::from_str(ok);
assert!(res.is_ok());
}
#[test]
fn test_rate_encoder_large_backlog_drains_exactly() {
let mut encoder = RateEncoder::try_new(0.0, 20_000_000.0, (0.0, 1.0), 1.0).unwrap();
let first = encoder.encode_step(&[1.0]);
assert_eq!(
first.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let second = encoder.encode_step(&[0.0]);
assert_eq!(
second.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let mut total = first.spikes.len() + second.spikes.len();
for _ in 0..10 {
total += encoder.encode_step(&[0.0]).spikes.len();
}
assert!(
total > RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP * 2,
"backlog should keep draining across steps, total={total}"
);
}
#[test]
fn test_rate_encoder_backlog_drains_above_f64_precision() {
#[cfg(feature = "serde")]
{
let seeded: RateEncoder = serde_json::from_str(
r#"{"base_rate":0.0,"max_rate":1.0,"range":[0.0,1.0],"dt_seconds":1.0,"accumulators":[2500.0]}"#,
)
.unwrap();
let mut encoder = seeded;
let first = encoder.encode_step(&[0.0]);
assert_eq!(
first.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let second = encoder.encode_step(&[0.0]);
assert_eq!(
second.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let third = encoder.encode_step(&[0.0]);
assert_eq!(third.spikes.len(), 452);
let fourth = encoder.encode_step(&[0.0]);
assert!(
fourth.spikes.is_empty(),
"backlog must fully drain rather than emit forever"
);
}
let mut encoder = RateEncoder::try_new(0.0, 1.0e16, (0.0, 1.0), 1.0).unwrap();
let first = encoder.encode_step(&[1.0]);
assert_eq!(
first.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let before = encoder.pending_spikes[0];
assert!(
before > RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP as u64,
"expected a large exact backlog, got {before}"
);
let quiet = encoder.encode_step(&[0.0]);
assert_eq!(
quiet.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
assert_eq!(
encoder.pending_spikes[0],
before - RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP as u64,
"u64 pending must decrement exactly past the f64 precision cliff"
);
}
#[test]
fn test_rate_encoder_streaming_bounds_extreme_dt() {
let mut encoder = RateEncoder::try_new(0.0, 10.0, (0.0, 1.0), f32::MAX).unwrap();
let output = encoder.encode_step(&[1.0]);
assert_eq!(
output.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let mut encoder = RateEncoder::try_new(0.0, 1.0e6, (0.0, 1.0), 1.0).unwrap();
let output = encoder.encode_step(&[1.0]);
assert_eq!(
output.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
let next = encoder.encode_step(&[0.0]);
assert_eq!(
next.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
}
#[test]
fn test_rate_encoder_nan_input_is_silent() {
let mut encoder = RateEncoder::try_new(0.0, 10.0, (0.0, 1.0), 0.1).unwrap();
let batch = encoder.encode(&[f32::NAN]);
assert!(
batch.spikes.is_empty(),
"NaN sensor values must not map to max-rate / p≈1"
);
let step = encoder.encode_step(&[f32::NAN]);
assert!(step.spikes.is_empty());
assert_eq!(
encoder.pending_spikes.first().copied().unwrap_or(0),
0,
"NaN must not seed a huge pending backlog"
);
let ok = encoder.encode_step(&[1.0]);
assert_eq!(ok.spikes.len(), 1);
}
#[test]
fn test_rate_encoder_rate_dt_product_saturates_not_silent() {
let mut encoder = RateEncoder::try_new(0.0, 1.0e38, (0.0, 1.0), 10.0).unwrap();
let first = encoder.encode_step(&[1.0]);
assert_eq!(
first.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP,
"rate×dt f32 overflow must not silence streaming"
);
assert!(
encoder.pending_spikes[0] > 0,
"expected a queued backlog after the per-step cap"
);
}
#[test]
fn test_rate_encoder_inactive_gain_clears_pending() {
let mut encoder = RateEncoder::try_new(0.0, 1.0e6, (0.0, 1.0), 1.0).unwrap();
let first = encoder.encode_step(&[1.0]);
assert_eq!(
first.spikes.len(),
RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP
);
assert!(encoder.pending_spikes[0] > 0);
let silenced = encoder.encode_step_with_rate_scale(&[0.0], 0.0);
assert!(silenced.spikes.is_empty());
assert_eq!(encoder.pending_spikes[0], 0);
assert_eq!(encoder.phases[0], 0.0);
let quiet = encoder.encode_step(&[0.0]);
assert!(quiet.spikes.is_empty());
}
#[test]
fn test_rate_encoder_default_dt_preserves_streaming_compatibility() {
let mut encoder = RateEncoder::new(0.0, 10.0, (0.0, 1.0));
assert_eq!(encoder.dt_seconds(), RateEncoder::DEFAULT_DT_SECONDS);
assert_eq!(encoder.encode_step(&[1.0]).spikes.len(), 1);
}
#[test]
fn test_rate_encoder_overflowed_gain_rate_saturates_not_silent() {
let mut encoder = RateEncoder::try_new(0.0, 1.0e35, (0.0, 1.0), 0.01).unwrap();
let gains = EncodingGains {
firing_rate_scale: 1.0e4,
..Default::default()
};
let mut spikes = 0usize;
for _ in 0..32 {
spikes += encoder.encode_with_gains(&[1.0], gains).spikes.len();
}
assert!(
spikes >= 28,
"overflowed modulated rate should saturate near p=1, got {spikes}/32 spikes"
);
}
#[test]
fn test_rate_encoder_streaming_uses_hz_times_dt() {
let cases = [(5.0, 0.2, 10), (20.0, 0.05, 20), (7.5, 0.1, 40)];
for (rate_hz, dt_seconds, steps) in cases {
let mut encoder = RateEncoder::try_new(0.0, rate_hz, (0.0, 1.0), dt_seconds).unwrap();
let spikes: usize = (0..steps)
.map(|_| encoder.encode_step(&[1.0]).spikes.len())
.sum();
let elapsed_seconds = dt_seconds * steps as f32;
let observed_hz = spikes as f32 / elapsed_seconds;
assert!(
(observed_hz - rate_hz).abs() <= 1.0 / elapsed_seconds,
"rate_hz={rate_hz}, dt={dt_seconds}, observed={observed_hz}"
);
}
}
#[test]
fn test_rate_encoder_stochastic_mean_matches_poisson_probability() {
let cases = [(2.0, 0.01), (10.0, 0.005), (25.0, 0.002)];
let trials = 50_000;
for (rate_hz, dt_seconds) in cases {
let mut encoder = RateEncoder::try_new(0.0, rate_hz, (0.0, 1.0), dt_seconds).unwrap();
let spikes: usize = (0..trials)
.map(|_| encoder.encode(&[1.0]).spikes.len())
.sum();
let observed_probability = spikes as f32 / trials as f32;
let expected_probability =
crate::poisson::probability_from_rate_hz(rate_hz, dt_seconds);
assert!(
(observed_probability - expected_probability).abs() < 0.01,
"rate_hz={rate_hz}, dt={dt_seconds}, observed_p={observed_probability}, expected_p={expected_probability}"
);
}
}
}
#[cfg(test)]
mod property_tests {
use super::*;
use crate::encoders::property_support::{
TRIALS, assert_unique_channel_spikes, sample_gain_scale, sample_input_value,
sample_positive_finite, scale_is_inactive,
};
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
const SEED: u64 = 0xAE69_0001;
fn sample_valid_encoder(rng: &mut StdRng) -> RateEncoder {
loop {
let base = sample_positive_finite(rng) * rng.random::<f32>();
let max = base + sample_positive_finite(rng);
let lo = rng.random_range(-50.0_f32..50.0);
let hi = lo + sample_positive_finite(rng);
let dt = sample_positive_finite(rng).clamp(1e-4, 1.0);
if let Ok(enc) = RateEncoder::try_new(base, max, (lo, hi), dt) {
return enc;
}
}
}
fn sample_input_vec(rng: &mut StdRng, n: usize, range: (f32, f32)) -> Vec<f32> {
(0..n).map(|_| sample_input_value(rng, range)).collect()
}
fn assert_both_silent(trial: usize, batch: &EncodedOutput, step: &EncodedOutput, why: &str) {
assert!(
batch.spikes.is_empty() && step.spikes.is_empty(),
"trial {trial}: {why}"
);
}
fn assert_active_batch_bounds(trial: usize, batch: &EncodedOutput, n_channels: usize) {
assert!(
batch.spikes.len() <= n_channels,
"trial {trial}: batch spikes {} > channels {n_channels}",
batch.spikes.len()
);
assert_unique_channel_spikes(&batch.spikes, n_channels);
}
fn assert_active_step_bounds(trial: usize, step: &EncodedOutput, n_channels: usize) {
let max_step = RateEncoder::MAX_SPIKES_PER_CHANNEL_PER_STEP.saturating_mul(n_channels);
assert!(
step.spikes.len() <= max_step,
"trial {trial}: step spikes {} exceed bound {max_step}",
step.spikes.len()
);
for spike in &step.spikes {
assert!((spike.channel as usize) < n_channels);
assert!(spike.polarity);
}
}
#[test]
fn prop_rate_silence_and_channel_bounds() {
let mut rng = StdRng::seed_from_u64(SEED);
for trial in 0..TRIALS {
let mut encoder = sample_valid_encoder(&mut rng);
let n = rng.random_range(0usize..=8);
let input = sample_input_vec(&mut rng, n, (0.0, 1.0));
let scale = sample_gain_scale(&mut rng);
let batch = encoder.encode_with_rate_scale(&input, scale);
let step = encoder.encode_step_with_rate_scale(&input, scale);
if input.is_empty() {
assert_both_silent(
trial,
&batch,
&step,
"empty input must silence batch and step",
);
continue;
}
if scale_is_inactive(scale) {
assert_both_silent(
trial,
&batch,
&step,
&format!("inactive rate_scale={scale:?} must silence"),
);
continue;
}
assert_active_batch_bounds(trial, &batch, input.len());
assert_active_step_bounds(trial, &step, input.len());
if input.iter().all(|v| !v.is_finite()) {
assert!(
batch.spikes.is_empty(),
"trial {trial}: all non-finite inputs must silence batch"
);
}
}
}
#[test]
fn prop_rate_probability_stays_in_unit_interval() {
let mut rng = StdRng::seed_from_u64(SEED ^ 0xB0B5);
for _ in 0..TRIALS {
let rate = match rng.random_range(0u8..8) {
0 => 0.0,
1 => -1.0,
2 => f32::NAN,
3 => f32::INFINITY,
4 => f32::MAX,
_ => sample_positive_finite(&mut rng) * rng.random_range(0.0_f32..100.0),
};
let dt = match rng.random_range(0u8..6) {
0 => 0.0,
1 => -0.1,
2 => f32::NAN,
3 => f32::INFINITY,
_ => sample_positive_finite(&mut rng).clamp(1e-6, 2.0),
};
let p = crate::poisson::probability_from_rate_hz(rate, dt);
assert!(
p.is_finite() && (0.0..=1.0).contains(&p),
"probability_from_rate_hz({rate}, {dt}) = {p}"
);
}
}
#[test]
fn prop_rate_encode_never_panics_on_sampled_inputs() {
let mut rng = StdRng::seed_from_u64(SEED ^ 0xBAD5);
for _ in 0..TRIALS {
let mut encoder = sample_valid_encoder(&mut rng);
let n = rng.random_range(0usize..=16);
let input = sample_input_vec(&mut rng, n, (-10.0, 10.0));
let scale = sample_gain_scale(&mut rng);
let _ = encoder.encode_with_rate_scale(&input, scale);
let _ = encoder.encode_step_with_rate_scale(&input, scale);
encoder.reset();
let _ = encoder.encode(&input);
let _ = encoder.encode_step(&input);
}
}
}