use crate::prelude::*;
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct PhaseEncoder {
cycle_steps: u64,
range: (f32, f32),
current_phase: u64,
}
fn validate_params(cycle_steps: u64, range: (f32, f32)) -> Result<(), EncoderError> {
if cycle_steps == 0 {
return Err(EncoderError::WindowMustBePositive {
parameter: "cycle_steps",
});
}
crate::error::validate_range("range", range)
}
impl PhaseEncoder {
pub fn new(cycle_steps: u64, range: (f32, f32)) -> Self {
Self::try_new(cycle_steps, range).unwrap_or_else(|error| panic!("{error}"))
}
pub fn try_new(cycle_steps: u64, range: (f32, f32)) -> Result<Self, EncoderError> {
validate_params(cycle_steps, range)?;
Ok(Self {
cycle_steps,
range,
current_phase: 0,
})
}
fn normalize(&self, value: f32) -> f64 {
let clamped = value.clamp(self.range.0, self.range.1) as f64;
let lo = self.range.0 as f64;
let hi = self.range.1 as f64;
(clamped - lo) / (hi - lo)
}
fn phase_offset(&self, normalized: f64) -> u64 {
((normalized * self.cycle_steps as f64).floor() as u64).min(self.cycle_steps - 1)
}
fn encode_current_cycle(&self, input: &[f32]) -> EncodedOutput {
let mut output = EncodedOutput::new();
for (channel, &value) in input.iter().enumerate() {
if !value.is_finite() {
continue;
}
let Ok(channel_u16) = u16::try_from(channel) else {
break;
};
let phase_offset = self.phase_offset(self.normalize(value));
output.spikes.push(SpikeEvent {
channel: channel_u16,
timestamp: self.current_phase.saturating_add(phase_offset),
polarity: true,
});
}
output
}
fn advance_phase(&mut self) {
self.current_phase = self.current_phase.saturating_add(1);
}
fn encode_current_cycle_with_sensitivity_scale(
&self,
input: &[f32],
sensitivity_scale: f32,
) -> EncodedOutput {
let mut output = EncodedOutput::new();
if !sensitivity_scale.is_finite() || sensitivity_scale <= 0.0 {
return output;
}
let lo = self.range.0 as f64;
let hi = lo + (self.range.1 as f64 - lo) * (sensitivity_scale as f64);
for (channel, &value) in input.iter().enumerate() {
if !value.is_finite() {
continue;
}
let Ok(channel_u16) = u16::try_from(channel) else {
break;
};
let normalized = ((value as f64 - lo) / (hi - lo)).clamp(0.0, 1.0);
let phase_offset = self.phase_offset(normalized);
output.spikes.push(SpikeEvent {
channel: channel_u16,
timestamp: self.current_phase.saturating_add(phase_offset),
polarity: true,
});
}
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,
)
}
}
impl Encoder for PhaseEncoder {
fn encode(&mut self, input: &[f32]) -> EncodedOutput {
let output = self.encode_current_cycle(input);
self.advance_phase();
output
}
fn encode_step(&mut self, input: &[f32]) -> EncodedOutput {
let output = self.encode_current_cycle(input);
self.advance_phase();
output
}
fn reset(&mut self) {
self.current_phase = 0;
}
}
impl ModulatedEncoder for PhaseEncoder {
fn encode_with_gains(&mut self, input: &[f32], gains: EncodingGains) -> EncodedOutput {
let output = self
.encode_current_cycle_with_sensitivity_scale(input, gains.sanitize().sensitivity_scale);
self.advance_phase();
output
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for PhaseEncoder {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Helper {
cycle_steps: u64,
range: (f32, f32),
#[serde(default)]
current_phase: u64,
}
let helper = Helper::deserialize(deserializer)?;
validate_params(helper.cycle_steps, helper.range).map_err(serde::de::Error::custom)?;
Ok(Self {
cycle_steps: helper.cycle_steps,
range: helper.range,
current_phase: helper.current_phase,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_wide_range_normalizes_without_nan() {
let mut encoder = PhaseEncoder::new(8, (f32::MIN, f32::MAX));
let output = encoder.encode(&[f32::MAX]);
assert_eq!(output.spikes.len(), 1);
assert_eq!(output.spikes[0].timestamp, 7);
}
#[test]
fn test_phase_mapping_clamps_and_quantizes() {
let mut encoder = PhaseEncoder::new(8, (0.0, 10.0));
let output = encoder.encode(&[-5.0, 0.0, 5.0, 10.0, 15.0]);
let timestamps: Vec<u64> = output.spikes.iter().map(|spike| spike.timestamp).collect();
let polarities: Vec<bool> = output.spikes.iter().map(|spike| spike.polarity).collect();
assert_eq!(timestamps, vec![0, 0, 4, 7, 7]);
assert_eq!(polarities, vec![true; 5]);
}
#[test]
fn test_phase_advances_after_each_call() {
let mut encoder = PhaseEncoder::new(4, (0.0, 1.0));
assert_eq!(encoder.encode(&[0.0]).spikes[0].timestamp, 0);
assert_eq!(encoder.encode(&[0.0]).spikes[0].timestamp, 1);
assert_eq!(encoder.encode_step(&[0.0]).spikes[0].timestamp, 2);
assert_eq!(encoder.encode_step(&[0.0]).spikes[0].timestamp, 3);
assert_eq!(encoder.encode(&[0.0]).spikes[0].timestamp, 4);
assert_eq!(encoder.encode(&[0.0]).spikes[0].timestamp % 4, 1);
}
#[test]
fn test_within_call_ordering_preserved_after_phase_advance() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
for _ in 0..6 {
encoder.encode(&[0.0]);
}
let output = encoder.encode(&[0.125, 0.375]); let timestamps: Vec<u64> = output.spikes.iter().map(|s| s.timestamp).collect();
assert_eq!(timestamps, vec![7, 9]);
assert!(timestamps[0] < timestamps[1]);
}
#[test]
fn test_reset_restores_initial_phase() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
encoder.encode(&[0.0]);
encoder.encode(&[0.0]);
encoder.reset();
let output = encoder.encode(&[1.0]);
assert_eq!(output.spikes[0].timestamp, 7);
}
#[test]
fn test_empty_input_returns_no_spikes() {
let mut encoder = PhaseEncoder::new(4, (0.0, 1.0));
let output = encoder.encode(&[]);
assert!(output.spikes.is_empty());
let next_output = encoder.encode(&[0.0]);
assert_eq!(next_output.spikes[0].timestamp, 1);
}
#[test]
fn test_nan_input_skips_channel() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
let output = encoder.encode(&[0.0, f32::NAN, 1.0]);
assert_eq!(output.spikes.len(), 2);
assert_eq!(output.spikes[0].channel, 0);
assert_eq!(output.spikes[1].channel, 2);
}
#[test]
#[should_panic(expected = "cycle_steps must be greater than 0")]
fn test_zero_cycle_steps_rejected() {
let _ = PhaseEncoder::new(0, (0.0, 1.0));
}
#[test]
#[should_panic(expected = "range must be finite and min must be less than max")]
fn test_invalid_range_rejected() {
let _ = PhaseEncoder::new(8, (1.0, 1.0));
}
#[test]
fn test_encode_step_matches_encode() {
let input = [2.5, 7.5];
let mut encode_encoder = PhaseEncoder::new(8, (0.0, 10.0));
let mut step_encoder = PhaseEncoder::new(8, (0.0, 10.0));
assert_eq!(
encode_encoder.encode(&input),
step_encoder.encode_step(&input)
);
assert_eq!(
encode_encoder.encode(&input),
step_encoder.encode_step(&input)
);
}
#[cfg(feature = "serde")]
#[test]
fn test_deserialize_rejects_zero_cycle_steps() {
let json = r#"{"cycle_steps":0,"range":[0.0,1.0],"current_phase":0}"#;
let err = serde_json::from_str::<PhaseEncoder>(json).unwrap_err();
assert!(err.to_string().contains("cycle_steps"));
}
#[test]
fn test_encode_with_modulators_identity() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
let curves = NeuromodulatorGainCurves::default();
let mods = NeuroModulators::default();
let plain = encoder.encode(&[0.5]);
let mut encoder2 = PhaseEncoder::new(8, (0.0, 1.0));
let modulated = encoder2.encode_with_modulators(&[0.5], &mods, &curves);
assert_eq!(plain.spikes[0].timestamp, modulated.spikes[0].timestamp);
}
#[test]
fn test_encode_with_modulators_sensitivity_scale() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
let curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
sensitivity: Some(GainCurve::new((0.0, 1.0), (0.5, 0.5))),
..Default::default()
},
..Default::default()
};
let mods = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let output = encoder.encode_with_modulators(&[0.5], &mods, &curves);
assert_eq!(output.spikes[0].timestamp, 7);
}
#[test]
fn test_encode_step_with_modulators_matches_encode() {
let input = [0.5];
let curves = NeuromodulatorGainCurves::default();
let mods = NeuroModulators::default();
let mut encoder1 = PhaseEncoder::new(8, (0.0, 1.0));
let mut encoder2 = PhaseEncoder::new(8, (0.0, 1.0));
let batch = encoder1.encode_with_modulators(&input, &mods, &curves);
let step = encoder2.encode_step_with_modulators(&input, &mods, &curves);
assert_eq!(batch, step);
}
#[test]
fn test_encode_with_modulators_zero_sensitivity_suppresses() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
let curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
sensitivity: Some(GainCurve::new((0.0, 1.0), (0.0, 0.0))),
..Default::default()
},
..Default::default()
};
let mods = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let output = encoder.encode_with_modulators(&[0.5], &mods, &curves);
assert!(output.spikes.is_empty());
}
#[test]
fn test_encode_with_modulators_nan_input_skips() {
let mut encoder = PhaseEncoder::new(8, (0.0, 1.0));
let curves = NeuromodulatorGainCurves {
dopamine: ModulatorGainCurves {
sensitivity: Some(GainCurve::new((0.0, 1.0), (1.0, 1.0))),
..Default::default()
},
..Default::default()
};
let mods = NeuroModulators {
dopamine: 1.0,
..Default::default()
};
let output = encoder.encode_with_modulators(&[0.0, f32::NAN, 1.0], &mods, &curves);
assert_eq!(output.spikes.len(), 2);
assert_eq!(output.spikes[0].channel, 0);
assert_eq!(output.spikes[1].channel, 2);
}
#[test]
fn test_phase_encoder_try_new_validation() {
assert_eq!(
PhaseEncoder::try_new(0, (0.0, 1.0)).err(),
Some(EncoderError::WindowMustBePositive {
parameter: "cycle_steps"
})
);
assert_eq!(
PhaseEncoder::try_new(1, (1.0, 1.0)).err(),
Some(EncoderError::InvalidRange { parameter: "range" })
);
}
}