use crate::prelude::*;
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct DeltaEncoder {
last_values: Vec<f32>,
threshold: f32,
}
impl DeltaEncoder {
pub fn new(threshold: f32, num_channels: usize) -> Self {
Self::try_new(threshold, num_channels).expect("invalid DeltaEncoder configuration")
}
pub fn try_new(threshold: f32, num_channels: usize) -> Result<Self, EncoderError> {
crate::error::validate_non_negative_finite("threshold", threshold)?;
crate::error::validate_channel_count(num_channels)?;
Ok(Self {
last_values: vec![0.0; num_channels],
threshold,
})
}
fn encode_with_threshold_scale(
&mut self,
input: &[f32],
threshold_scale: f32,
) -> EncodedOutput {
let mut output = EncodedOutput::new();
let effective_threshold = (self.threshold * threshold_scale).max(0.0);
for (i, &value) in input.iter().enumerate() {
if i >= self.last_values.len() {
break;
}
let Ok(channel) = u16::try_from(i) else {
break;
};
let delta = (value - self.last_values[i]).abs();
if delta > effective_threshold {
output.spikes.push(SpikeEvent {
channel,
timestamp: 0,
polarity: value > self.last_values[i],
});
self.last_values[i] = value;
}
}
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 DeltaEncoder {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Helper {
last_values: Vec<f32>,
threshold: f32,
}
let helper = Helper::deserialize(deserializer)?;
let mut encoder = Self::try_new(helper.threshold, helper.last_values.len())
.map_err(serde::de::Error::custom)?;
if helper.last_values.iter().any(|value| !value.is_finite()) {
return Err(serde::de::Error::custom("last_values must be finite"));
}
encoder.last_values = helper.last_values;
Ok(encoder)
}
}
impl Encoder for DeltaEncoder {
fn encode(&mut self, input: &[f32]) -> EncodedOutput {
self.encode_with_threshold_scale(input, 1.0)
}
fn encode_step(&mut self, input: &[f32]) -> EncodedOutput {
let safe_input = if input.len() > self.last_values.len() {
&input[..self.last_values.len()]
} else {
input
};
self.encode(safe_input)
}
fn reset(&mut self) {
for val in self.last_values.iter_mut() {
*val = 0.0;
}
}
}
impl ModulatedEncoder for DeltaEncoder {
fn encode_with_gains(&mut self, input: &[f32], gains: EncodingGains) -> EncodedOutput {
self.encode_with_threshold_scale(input, gains.sanitize().threshold_scale)
}
}
pub fn encode_deltas_to_spikes(deltas: &[f32], threshold: f32) -> Vec<bool> {
deltas.iter().map(|&d| d.abs() > threshold).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_delta_encoder() {
let mut encoder = DeltaEncoder::new(2.0, 1);
let output = encoder.encode(&[1.0]); assert!(output.spikes.is_empty());
let output = encoder.encode(&[3.5]); assert!(!output.spikes.is_empty());
assert!(output.spikes[0].polarity);
let output = encoder.encode(&[4.0]); assert!(output.spikes.is_empty());
let output = encoder.encode(&[1.0]); assert!(!output.spikes.is_empty());
assert!(!output.spikes[0].polarity);
}
#[test]
fn test_delta_encoder_encode_step() {
let mut encoder = DeltaEncoder::new(2.0, 2);
let output = encoder.encode_step(&[3.0, 3.0, 3.0]); assert_eq!(output.spikes.len(), 2);
}
#[test]
fn test_delta_encoder_multi_channel_reset() {
let mut encoder = DeltaEncoder::new(1.0, 2);
encoder.encode(&[2.0, 2.0]);
assert_eq!(encoder.last_values, vec![2.0, 2.0]);
encoder.reset();
assert_eq!(encoder.last_values, vec![0.0, 0.0]);
}
#[test]
fn test_delta_encoder_empty_input() {
let mut encoder = DeltaEncoder::new(1.0, 5);
let output = encoder.encode(&[]);
assert!(output.spikes.is_empty());
}
#[test]
fn test_encode_deltas_to_spikes() {
let deltas = [0.1, 0.5, -0.8, 1.2];
let threshold = 0.7;
let spikes = encode_deltas_to_spikes(&deltas, threshold);
assert_eq!(spikes, vec![false, false, true, true]);
}
#[test]
fn test_delta_encoder_modulators_reduce_threshold() {
let mut encoder = DeltaEncoder::new(1.0, 1);
let modulators = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let gain_curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
threshold: Some(GainCurve::new((0.0, 1.0), (1.0, 0.5))),
..Default::default()
},
..Default::default()
};
assert!(encoder.encode(&[0.75]).spikes.is_empty());
encoder.reset();
let modulated = encoder.encode_with_modulators(&[0.75], &modulators, &gain_curves);
assert_eq!(modulated.spikes.len(), 1);
}
#[test]
fn test_delta_encoder_encode_step_with_modulators() {
let mut encoder = DeltaEncoder::new(1.0, 1);
let modulators = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let gain_curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
threshold: Some(GainCurve::new((0.0, 1.0), (1.0, 0.5))),
..Default::default()
},
..Default::default()
};
assert!(
encoder
.encode_step_with_modulators(&[0.0], &modulators, &gain_curves)
.spikes
.is_empty()
);
let modulated = encoder.encode_step_with_modulators(&[0.75], &modulators, &gain_curves);
assert_eq!(modulated.spikes.len(), 1);
}
#[test]
fn test_delta_encoder_step_shorter_input() {
let mut encoder = DeltaEncoder::new(1.0, 2);
let output = encoder.encode_step(&[2.0]);
assert_eq!(output.spikes.len(), 1);
}
#[test]
fn test_delta_encoder_truncates_excess_channels() {
let mut encoder = DeltaEncoder::new(1.0, 1);
let output = encoder.encode(&[2.0, 3.0]);
assert_eq!(output.spikes.len(), 1);
}
#[test]
fn test_delta_encoder_zero_threshold_scale_spikes_on_any_change() {
let mut encoder = DeltaEncoder::new(1.0, 1);
encoder.encode(&[0.0]);
let output = encoder.encode_with_threshold_scale(&[0.01], 0.0);
assert_eq!(output.spikes.len(), 1);
}
#[test]
fn test_delta_encoder_zero_threshold_spikes_on_any_change() {
let mut encoder = DeltaEncoder::new(0.0, 1);
encoder.encode(&[0.0]);
let output = encoder.encode(&[0.01]);
assert_eq!(output.spikes.len(), 1);
let quiet = encoder.encode(&[0.01]);
assert!(quiet.spikes.is_empty());
}
#[test]
fn test_delta_encoder_try_new_validation() {
assert!(DeltaEncoder::try_new(0.0, 1).is_ok());
assert_eq!(
DeltaEncoder::try_new(-1.0, 1).err(),
Some(EncoderError::NonNegativeFinite {
parameter: "threshold"
})
);
assert_eq!(
DeltaEncoder::try_new(f32::NAN, 1).err(),
Some(EncoderError::NonNegativeFinite {
parameter: "threshold"
})
);
assert_eq!(
DeltaEncoder::try_new(1.0, u16::MAX as usize + 2).err(),
Some(EncoderError::NumChannelsTooLarge)
);
}
}