Skip to main content

axon_encoder/
modulators.rs

1//! Encoding-side gain controls and optional named modulator *bags*.
2//!
3//! All types here live in **this** crate only (no external neuromodulator
4//! runtime dependency):
5//!
6//! - [`EncodingGains`] — generic scales (rate / threshold / latency / sensitivity).
7//! - [`NeuroModulators`] / [`NeuromodulatorGainCurves`] — encoding-local helpers
8//!   with biologically familiar field names for evaluating gains. Map your own
9//!   application state into [`EncodingGains`] (see
10//!   `examples/sibling_gains_adapter.rs` for a pattern).
11
12const EVENT_DOPAMINE_DECAY: f32 = 0.95;
13const CORTISOL_DECAY: f32 = 0.90;
14const ACETYLCHOLINE_DECAY: f32 = 0.99;
15const TEMPO_DECAY: f32 = 0.98;
16/// Allow true zero gain (full silence / zero threshold). Non-finite values map
17/// to identity; values above this cap are clamped for numerical stability.
18const MIN_GAIN_SCALE: f32 = 0.0;
19const MAX_GAIN_SCALE: f32 = 1e4;
20
21fn sanitize_gain_scale(scale: f32) -> f32 {
22    if !scale.is_finite() {
23        return 1.0;
24    }
25
26    scale.clamp(MIN_GAIN_SCALE, MAX_GAIN_SCALE)
27}
28
29/// Neuromodulator levels consumed by gain curves.
30///
31/// Fields are public `f32` values with no constructor validation. Callers
32/// typically keep levels ≥ 0; negative values are not rejected here. When a
33/// level is fed through a [`GainCurve`], it is clamped to that curve's input
34/// range before interpolation. Call [`NeuroModulators::decay`] between steps
35/// for the fixed exponential decay schedule (decay floors at 0).
36///
37/// # Examples
38///
39/// ```rust
40/// use axon_encoder::prelude::*;
41///
42/// let mut mods = NeuroModulators {
43///     dopamine: 1.0,
44///     ..Default::default()
45/// };
46/// mods.decay();
47/// assert!(mods.dopamine < 1.0);
48/// assert!(mods.dopamine >= 0.0);
49/// ```
50#[derive(Debug, Clone, Copy, Default, PartialEq)]
51#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
52pub struct NeuroModulators {
53    pub dopamine: f32,
54    pub cortisol: f32,
55    pub acetylcholine: f32,
56    pub tempo: f32,
57}
58
59impl NeuroModulators {
60    pub fn decay(&mut self) {
61        self.dopamine = (self.dopamine * EVENT_DOPAMINE_DECAY).max(0.0);
62        self.cortisol = (self.cortisol * CORTISOL_DECAY).max(0.0);
63        self.acetylcholine = (self.acetylcholine * ACETYLCHOLINE_DECAY).max(0.0);
64        self.tempo = (self.tempo * TEMPO_DECAY).max(0.0);
65    }
66}
67
68/// Piecewise-linear map from a modulator level to a gain scale.
69///
70/// # Examples
71///
72/// ```rust
73/// use axon_encoder::prelude::*;
74///
75/// // Map level 0..1 to gain 1..2 (identity at mid-point is 1.5).
76/// let curve = GainCurve::new((0.0, 1.0), (1.0, 2.0));
77/// assert!((curve.evaluate(0.5) - 1.5).abs() < 1e-5);
78/// ```
79#[derive(Debug, Clone, Copy, PartialEq)]
80#[cfg_attr(feature = "serde", derive(serde::Serialize))]
81pub struct GainCurve {
82    pub input_range: (f32, f32),
83    pub output_range: (f32, f32),
84}
85
86impl GainCurve {
87    pub fn new(input_range: (f32, f32), output_range: (f32, f32)) -> Self {
88        assert!(
89            input_range.0.is_finite() && input_range.1.is_finite() && input_range.0 < input_range.1,
90            "input_range min must be less than max and finite"
91        );
92        assert!(
93            output_range.0.is_finite() && output_range.1.is_finite(),
94            "output_range values must be finite"
95        );
96
97        Self {
98            input_range,
99            output_range,
100        }
101    }
102
103    pub fn identity() -> Self {
104        Self {
105            input_range: (0.0, 1.0),
106            output_range: (1.0, 1.0),
107        }
108    }
109
110    /// Returns whether this curve has a valid, finite, ordered input range.
111    fn has_valid_input_range(&self) -> bool {
112        self.input_range.0.is_finite()
113            && self.input_range.1.is_finite()
114            && self.input_range.0 < self.input_range.1
115    }
116
117    /// Evaluate the gain curve at the given modulator level.
118    ///
119    /// Negative levels are clamped to `input_range.0`. NaN or non-finite
120    /// levels return the identity gain (1.0).
121    pub fn evaluate(&self, level: f32) -> f32 {
122        // Guard against NaN levels and invalid ranges that can arise from
123        // public fields or bypassed constructors (e.g. deserialization).
124        if !level.is_finite()
125            || !self.has_valid_input_range()
126            || !self.output_range.0.is_finite()
127            || !self.output_range.1.is_finite()
128        {
129            return 1.0;
130        }
131
132        let clamped_level = level.clamp(self.input_range.0, self.input_range.1);
133        // Use f64 for span to avoid overflow for valid f32 ranges (e.g., f32::MIN..f32::MAX).
134        let span = (self.input_range.1 as f64) - (self.input_range.0 as f64);
135        // span is guaranteed > 0 by has_valid_input_range
136        let position = ((clamped_level as f64 - self.input_range.0 as f64) / span) as f32;
137
138        // Use lerp form to avoid overflow when output_range spans nearly f32::MAX.
139        let raw_scale = self.output_range.0 * (1.0 - position) + self.output_range.1 * position;
140
141        sanitize_gain_scale(raw_scale)
142    }
143}
144
145#[cfg(feature = "serde")]
146impl<'de> serde::Deserialize<'de> for GainCurve {
147    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
148    where
149        D: serde::Deserializer<'de>,
150    {
151        #[derive(serde::Deserialize)]
152        struct Helper {
153            input_range: (f32, f32),
154            output_range: (f32, f32),
155        }
156
157        let helper = Helper::deserialize(deserializer)?;
158
159        if !helper.input_range.0.is_finite()
160            || !helper.input_range.1.is_finite()
161            || helper.input_range.0 >= helper.input_range.1
162        {
163            return Err(serde::de::Error::custom(
164                "input_range min must be less than max and finite",
165            ));
166        }
167        if !helper.output_range.0.is_finite() || !helper.output_range.1.is_finite() {
168            return Err(serde::de::Error::custom(
169                "output_range values must be finite",
170            ));
171        }
172
173        Ok(Self {
174            input_range: helper.input_range,
175            output_range: helper.output_range,
176        })
177    }
178}
179
180impl Default for GainCurve {
181    fn default() -> Self {
182        Self::identity()
183    }
184}
185
186#[derive(Debug, Clone, Copy, PartialEq, Default)]
187#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
188pub struct ModulatorGainCurves {
189    pub threshold: Option<GainCurve>,
190    pub sensitivity: Option<GainCurve>,
191    pub firing_rate: Option<GainCurve>,
192    pub latency: Option<GainCurve>,
193}
194
195/// Scales produced by neuromodulator gain curves for each encoder component.
196///
197/// # Zero-gain semantics
198///
199/// The meaning of a 0.0 gain depends on the component:
200/// - `threshold_scale = 0.0` → effective threshold is 0 → every input spikes (maximum sensitivity)
201/// - `sensitivity_scale = 0.0` → output is suppressed (no spikes for PopulationEncoder)
202/// - `firing_rate_scale = 0.0` → firing rate is 0 → no spikes (silence)
203/// - `latency_scale = 0.0` → max_latency is 0 → all spikes at timestamp 0 (instant response)
204///
205/// This asymmetry is intentional and reflects the physical semantics of each component.
206///
207/// # Examples
208///
209/// ```rust
210/// use axon_encoder::prelude::*;
211///
212/// let gains = EncodingGains {
213///     firing_rate_scale: 0.0, // silence rate-based paths
214///     ..EncodingGains::identity()
215/// }
216/// .sanitize();
217/// assert_eq!(gains.firing_rate_scale, 0.0);
218/// ```
219#[derive(Debug, Clone, Copy, PartialEq)]
220#[cfg_attr(feature = "serde", derive(serde::Serialize))]
221pub struct EncodingGains {
222    pub threshold_scale: f32,
223    pub sensitivity_scale: f32,
224    pub firing_rate_scale: f32,
225    pub latency_scale: f32,
226}
227
228impl EncodingGains {
229    pub fn identity() -> Self {
230        Self {
231            threshold_scale: 1.0,
232            sensitivity_scale: 1.0,
233            firing_rate_scale: 1.0,
234            latency_scale: 1.0,
235        }
236    }
237
238    /// Clamps non-finite and out-of-range gain components to safe defaults.
239    pub fn sanitize(self) -> Self {
240        Self {
241            threshold_scale: sanitize_gain_scale(self.threshold_scale),
242            sensitivity_scale: sanitize_gain_scale(self.sensitivity_scale),
243            firing_rate_scale: sanitize_gain_scale(self.firing_rate_scale),
244            latency_scale: sanitize_gain_scale(self.latency_scale),
245        }
246    }
247}
248
249#[cfg(feature = "serde")]
250impl<'de> serde::Deserialize<'de> for EncodingGains {
251    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
252    where
253        D: serde::Deserializer<'de>,
254    {
255        #[derive(serde::Deserialize)]
256        struct Helper {
257            #[serde(default = "default_gain_scale")]
258            threshold_scale: f32,
259            #[serde(default = "default_gain_scale")]
260            sensitivity_scale: f32,
261            #[serde(default = "default_gain_scale")]
262            firing_rate_scale: f32,
263            #[serde(default = "default_gain_scale")]
264            latency_scale: f32,
265        }
266
267        fn default_gain_scale() -> f32 {
268            1.0
269        }
270
271        let helper = Helper::deserialize(deserializer)?;
272        let gains = Self {
273            threshold_scale: helper.threshold_scale,
274            sensitivity_scale: helper.sensitivity_scale,
275            firing_rate_scale: helper.firing_rate_scale,
276            latency_scale: helper.latency_scale,
277        };
278        Ok(gains.sanitize())
279    }
280}
281
282impl Default for EncodingGains {
283    fn default() -> Self {
284        Self::identity()
285    }
286}
287
288/// Per-neuromodulator gain curves composing into [`EncodingGains`].
289///
290/// Default curves are all identity (no modulation).
291///
292/// # Examples
293///
294/// ```rust
295/// use axon_encoder::prelude::*;
296///
297/// let curves = NeuromodulatorGainCurves {
298///     tempo: ModulatorGainCurves {
299///         sensitivity: Some(GainCurve::new((0.0, 1.0), (1.0, 2.0))),
300///         ..Default::default()
301///     },
302///     ..Default::default()
303/// };
304/// let mods = NeuroModulators {
305///     tempo: 1.0,
306///     ..Default::default()
307/// };
308/// let gains = curves.evaluate(&mods);
309/// assert!(gains.sensitivity_scale > 1.0);
310/// ```
311#[derive(Debug, Clone, Copy, PartialEq, Default)]
312#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
313pub struct NeuromodulatorGainCurves {
314    #[cfg_attr(feature = "serde", serde(default))]
315    pub dopamine: ModulatorGainCurves,
316    #[cfg_attr(feature = "serde", serde(default))]
317    pub cortisol: ModulatorGainCurves,
318    #[cfg_attr(feature = "serde", serde(default))]
319    pub acetylcholine: ModulatorGainCurves,
320    #[cfg_attr(feature = "serde", serde(default))]
321    pub tempo: ModulatorGainCurves,
322}
323
324impl NeuromodulatorGainCurves {
325    pub fn evaluate(&self, modulators: &NeuroModulators) -> EncodingGains {
326        let mut gains = EncodingGains::identity();
327
328        Self::apply_curves(&mut gains, self.dopamine, modulators.dopamine);
329        Self::apply_curves(&mut gains, self.cortisol, modulators.cortisol);
330        Self::apply_curves(&mut gains, self.acetylcholine, modulators.acetylcholine);
331        Self::apply_curves(&mut gains, self.tempo, modulators.tempo);
332
333        gains.sanitize()
334    }
335
336    fn apply_curves(gains: &mut EncodingGains, curves: ModulatorGainCurves, level: f32) {
337        if let Some(curve) = curves.threshold {
338            gains.threshold_scale *= curve.evaluate(level);
339        }
340        if let Some(curve) = curves.sensitivity {
341            gains.sensitivity_scale *= curve.evaluate(level);
342        }
343        if let Some(curve) = curves.firing_rate {
344            gains.firing_rate_scale *= curve.evaluate(level);
345        }
346        if let Some(curve) = curves.latency {
347            gains.latency_scale *= curve.evaluate(level);
348        }
349    }
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355
356    #[test]
357    fn gain_curve_clamps_input_range() {
358        let curve = GainCurve::new((0.0, 1.0), (0.5, 2.0));
359
360        assert_eq!(curve.evaluate(-5.0), 0.5);
361        assert_eq!(curve.evaluate(5.0), 2.0);
362    }
363
364    #[test]
365    fn gain_curve_interpolates_wide_f32_range() {
366        let curve = GainCurve::new((f32::MIN, f32::MAX), (0.0, 2.0));
367
368        assert_eq!(curve.evaluate(f32::MIN), 0.0);
369        assert_eq!(curve.evaluate(f32::MAX), 2.0);
370        assert!((curve.evaluate(0.0) - 1.0).abs() < 1e-5);
371    }
372
373    #[test]
374    fn gain_curve_sanitizes_invalid_outputs() {
375        let curve = GainCurve::new((0.0, 1.0), (-2.0, 2.0));
376
377        assert_eq!(curve.evaluate(0.0), MIN_GAIN_SCALE);
378        assert_eq!(curve.evaluate(f32::NAN), 1.0);
379    }
380
381    #[test]
382    fn gain_curve_allows_true_zero_output() {
383        let curve = GainCurve::new((0.0, 1.0), (0.0, 1.0));
384        assert_eq!(curve.evaluate(0.0), 0.0);
385    }
386
387    #[test]
388    fn gain_curve_invalid_range_returns_identity() {
389        // Bypass constructor the same way a bad public-field mutation would.
390        let curve = GainCurve {
391            input_range: (1.0, 1.0),
392            output_range: (0.0, 2.0),
393        };
394        assert_eq!(curve.evaluate(0.5), 1.0);
395    }
396
397    #[test]
398    fn neuromodulator_curves_compose_multiplicatively() {
399        let curves = NeuromodulatorGainCurves {
400            dopamine: ModulatorGainCurves {
401                firing_rate: Some(GainCurve::new((0.0, 1.0), (1.0, 2.0))),
402                ..Default::default()
403            },
404            cortisol: ModulatorGainCurves {
405                threshold: Some(GainCurve::new((0.0, 1.0), (1.0, 0.5))),
406                ..Default::default()
407            },
408            acetylcholine: ModulatorGainCurves {
409                firing_rate: Some(GainCurve::new((0.0, 1.0), (1.0, 1.5))),
410                ..Default::default()
411            },
412            tempo: ModulatorGainCurves {
413                sensitivity: Some(GainCurve::new((0.0, 1.0), (1.0, 1.25))),
414                ..Default::default()
415            },
416        };
417        let modulators = NeuroModulators {
418            dopamine: 1.0,
419            cortisol: 1.0,
420            acetylcholine: 1.0,
421            tempo: 1.0,
422        };
423
424        let gains = curves.evaluate(&modulators);
425
426        assert_eq!(gains.threshold_scale, 0.5);
427        assert_eq!(gains.sensitivity_scale, 1.25);
428        assert_eq!(gains.firing_rate_scale, 3.0);
429        assert_eq!(gains.latency_scale, 1.0); // no latency curve set
430    }
431
432    #[cfg(feature = "serde")]
433    #[test]
434    fn gain_curve_rejects_invalid_deserialize() {
435        let json = r#"{"input_range":[1.0,0.0],"output_range":[0.0,1.0]}"#;
436        let err = serde_json::from_str::<GainCurve>(json).unwrap_err();
437        assert!(err.to_string().contains("input_range"));
438    }
439
440    #[cfg(feature = "serde")]
441    #[test]
442    fn encoding_gains_partial_json_deserializes() {
443        let json = r#"{"threshold_scale":0.5}"#;
444        let gains: EncodingGains = serde_json::from_str(json).unwrap();
445        assert_eq!(gains.threshold_scale, 0.5);
446        assert_eq!(gains.sensitivity_scale, 1.0);
447        assert_eq!(gains.firing_rate_scale, 1.0);
448        assert_eq!(gains.latency_scale, 1.0);
449    }
450
451    #[cfg(feature = "serde")]
452    #[test]
453    fn encoding_gains_deserialize_sanitizes_values() {
454        // Use out-of-range values that serde_json can parse (NaN is not valid JSON)
455        let json =
456            r#"{"threshold_scale":-999.0,"sensitivity_scale":999999.0,"firing_rate_scale":0.5}"#;
457        let gains: EncodingGains = serde_json::from_str(json).unwrap();
458        assert_eq!(gains.threshold_scale, 0.0); // -999 clamped to MIN_GAIN_SCALE (0.0)
459        assert_eq!(gains.sensitivity_scale, MAX_GAIN_SCALE); // 999999 clamped to MAX_GAIN_SCALE
460        assert_eq!(gains.firing_rate_scale, 0.5); // in range, unchanged
461        assert_eq!(gains.latency_scale, 1.0); // defaults to 1.0 when omitted
462    }
463
464    #[test]
465    fn sanitize_gain_scale_handles_nan_and_infinity() {
466        assert_eq!(sanitize_gain_scale(f32::NAN), 1.0);
467        assert_eq!(sanitize_gain_scale(f32::INFINITY), 1.0);
468        assert_eq!(sanitize_gain_scale(f32::NEG_INFINITY), 1.0);
469        assert_eq!(sanitize_gain_scale(0.0), 0.0);
470        assert_eq!(sanitize_gain_scale(5.0), 5.0);
471        assert_eq!(sanitize_gain_scale(1e10), MAX_GAIN_SCALE);
472    }
473
474    #[test]
475    fn neuro_modulators_decay() {
476        let mut mods = NeuroModulators {
477            dopamine: 1.0,
478            cortisol: 1.0,
479            acetylcholine: 1.0,
480            tempo: 1.0,
481        };
482        mods.decay();
483        assert!((mods.dopamine - 0.95).abs() < 1e-6);
484        assert!((mods.cortisol - 0.90).abs() < 1e-6);
485        assert!((mods.acetylcholine - 0.99).abs() < 1e-6);
486        assert!((mods.tempo - 0.98).abs() < 1e-6);
487
488        // Decay floors at zero
489        mods.dopamine = -0.5;
490        mods.decay();
491        assert_eq!(mods.dopamine, 0.0);
492    }
493
494    #[test]
495    fn gain_curve_identity_returns_constant_one() {
496        let curve = GainCurve::identity();
497        assert_eq!(curve.evaluate(0.0), 1.0);
498        assert_eq!(curve.evaluate(0.5), 1.0);
499        assert_eq!(curve.evaluate(1.0), 1.0);
500    }
501
502    #[test]
503    fn gain_curve_evaluate_non_finite_output_range_returns_identity() {
504        let curve = GainCurve {
505            input_range: (0.0, 1.0),
506            output_range: (f32::NAN, 2.0),
507        };
508        assert_eq!(curve.evaluate(0.5), 1.0);
509
510        let curve2 = GainCurve {
511            input_range: (0.0, 1.0),
512            output_range: (1.0, f32::INFINITY),
513        };
514        assert_eq!(curve2.evaluate(0.5), 1.0);
515    }
516
517    #[test]
518    fn encoding_gains_sanitize_clamps_extremes() {
519        let gains = EncodingGains {
520            threshold_scale: f32::NAN,
521            sensitivity_scale: f32::INFINITY,
522            firing_rate_scale: -1.0,
523            latency_scale: 0.5,
524        };
525        let sanitized = gains.sanitize();
526        assert_eq!(sanitized.threshold_scale, 1.0);
527        assert_eq!(sanitized.sensitivity_scale, 1.0);
528        assert_eq!(sanitized.firing_rate_scale, 0.0);
529        assert_eq!(sanitized.latency_scale, 0.5);
530    }
531
532    #[test]
533    fn neuromodulator_curves_all_none_returns_identity() {
534        let curves = NeuromodulatorGainCurves::default();
535        let mods = NeuroModulators::default();
536        let gains = curves.evaluate(&mods);
537        assert_eq!(gains.threshold_scale, 1.0);
538        assert_eq!(gains.sensitivity_scale, 1.0);
539        assert_eq!(gains.firing_rate_scale, 1.0);
540        assert_eq!(gains.latency_scale, 1.0);
541    }
542
543    #[test]
544    fn neuromodulator_curves_partial_none() {
545        let curves = NeuromodulatorGainCurves {
546            dopamine: ModulatorGainCurves {
547                threshold: Some(GainCurve::new((0.0, 1.0), (1.0, 2.0))),
548                ..Default::default()
549            },
550            ..Default::default()
551        };
552        let mods = NeuroModulators {
553            dopamine: 1.0,
554            ..Default::default()
555        };
556        let gains = curves.evaluate(&mods);
557        assert_eq!(gains.threshold_scale, 2.0);
558        assert_eq!(gains.sensitivity_scale, 1.0);
559        assert_eq!(gains.firing_rate_scale, 1.0);
560        assert_eq!(gains.latency_scale, 1.0);
561    }
562
563    #[test]
564    fn modulator_gain_curves_default_is_none() {
565        let curves = ModulatorGainCurves::default();
566        assert!(curves.threshold.is_none());
567        assert!(curves.sensitivity.is_none());
568        assert!(curves.firing_rate.is_none());
569    }
570
571    #[test]
572    fn gain_curve_default_is_identity() {
573        assert_eq!(GainCurve::default(), GainCurve::identity());
574    }
575
576    #[test]
577    fn encoding_gains_default_is_identity() {
578        assert_eq!(EncodingGains::default(), EncodingGains::identity());
579    }
580
581    #[cfg(feature = "serde")]
582    #[test]
583    fn neuromodulator_gain_curves_partial_json_deserializes() {
584        // Only set dopamine; cortisol/acetylcholine/tempo should default
585        let json = r#"{
586            "dopamine": {
587                "firing_rate": {"input_range": [0.0, 1.0], "output_range": [1.0, 2.0]}
588            }
589        }"#;
590        let curves: NeuromodulatorGainCurves = serde_json::from_str(json).unwrap();
591        assert!(curves.dopamine.firing_rate.is_some());
592        assert!(curves.cortisol.threshold.is_none());
593        assert!(curves.acetylcholine.sensitivity.is_none());
594        assert!(curves.tempo.firing_rate.is_none());
595    }
596}