feagi_structures/genomic/
modulator.rs1use serde::{Deserialize, Serialize};
11use std::fmt;
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
15pub enum ModulatorKind {
16 FiringThreshold,
18 Leak,
20 FiringProbability,
22 TransmissionGain,
24 Reward,
26 LearningRate,
28}
29
30impl ModulatorKind {
31 pub fn as_str(self) -> &'static str {
33 match self {
34 Self::FiringThreshold => "neuro.firing_threshold",
35 Self::Leak => "neuro.leak",
36 Self::FiringProbability => "neuro.firing_probability",
37 Self::TransmissionGain => "synaptic.transmission_gain",
38 Self::Reward => "synaptic.reward",
39 Self::LearningRate => "synaptic.learning_rate",
40 }
41 }
42
43 pub fn parse(value: &str) -> Result<Self, String> {
45 match value {
46 "neuro.firing_threshold" => Ok(Self::FiringThreshold),
47 "neuro.leak" => Ok(Self::Leak),
48 "neuro.firing_probability" => Ok(Self::FiringProbability),
49 "synaptic.transmission_gain" => Ok(Self::TransmissionGain),
50 "synaptic.reward" => Ok(Self::Reward),
51 "synaptic.learning_rate" => Ok(Self::LearningRate),
52 other => Err(format!("unknown modulator type '{other}'")),
53 }
54 }
55
56 pub fn all() -> &'static [Self] {
58 &[
59 Self::FiringThreshold,
60 Self::Leak,
61 Self::FiringProbability,
62 Self::TransmissionGain,
63 Self::Reward,
64 Self::LearningRate,
65 ]
66 }
67
68 pub fn is_neuromodulator(self) -> bool {
70 matches!(
71 self,
72 Self::FiringThreshold | Self::Leak | Self::FiringProbability
73 )
74 }
75
76 pub fn is_synaptic(self) -> bool {
78 !self.is_neuromodulator()
79 }
80
81 pub fn combines_by_addition(self) -> bool {
83 matches!(self, Self::Reward)
84 }
85}
86
87impl fmt::Display for ModulatorKind {
88 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
89 f.write_str(self.as_str())
90 }
91}
92
93#[derive(Debug, Clone, PartialEq, Eq)]
95pub struct ModulatorValidationError(pub String);
96
97impl fmt::Display for ModulatorValidationError {
98 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
99 f.write_str(&self.0)
100 }
101}
102
103pub fn validate_instance_fields(
109 magnitude_percent: f32,
110 effect_duration_bursts: u16,
111 graded: bool,
112 full_scale_potential: Option<f32>,
113) -> Result<(), ModulatorValidationError> {
114 if !magnitude_percent.is_finite() || magnitude_percent < -100.0 {
115 return Err(ModulatorValidationError(
116 "magnitude_percent must be finite and >= -100".to_string(),
117 ));
118 }
119 if effect_duration_bursts < 1 {
120 return Err(ModulatorValidationError(
121 "effect_duration_bursts must be >= 1".to_string(),
122 ));
123 }
124 if graded {
125 match full_scale_potential {
126 Some(scale) if scale.is_finite() && scale > 0.0 => {}
127 _ => {
128 return Err(ModulatorValidationError(
129 "graded modulators require full_scale_potential > 0".to_string(),
130 ));
131 }
132 }
133 }
134 Ok(())
135}
136
137pub fn validate_spike_train(
139 enabled: bool,
140 consecutive_fire_limit: u16,
141) -> Result<(), ModulatorValidationError> {
142 if enabled && consecutive_fire_limit < 1 {
143 return Err(ModulatorValidationError(
144 "spike_train requires consecutive_fire_limit >= 1".to_string(),
145 ));
146 }
147 Ok(())
148}
149
150pub fn modulator_signal(
156 magnitude_percent: f32,
157 fired: bool,
158 graded: bool,
159 membrane_potential: f32,
160 full_scale_potential: f32,
161) -> f32 {
162 if !fired {
163 return 0.0;
164 }
165 let strength = if graded {
166 (membrane_potential / full_scale_potential).min(1.0)
167 } else {
168 1.0
169 };
170 (magnitude_percent / 100.0) * strength
171}
172
173pub fn multiplicative_factor(signals: &[f32]) -> f32 {
175 signals
176 .iter()
177 .fold(1.0_f32, |acc, signal| acc * (1.0 + signal))
178}
179
180pub fn summed_reward(signals: &[f32]) -> f32 {
182 signals.iter().sum()
183}
184
185pub fn scale_baseline(kind: ModulatorKind, baseline: f32, factor: f32) -> f32 {
190 let value = baseline * factor;
191 match kind {
192 ModulatorKind::Leak | ModulatorKind::FiringProbability => value.clamp(0.0, 1.0),
193 ModulatorKind::FiringThreshold
194 | ModulatorKind::TransmissionGain
195 | ModulatorKind::LearningRate => value.max(0.0),
196 ModulatorKind::Reward => value,
197 }
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203
204 #[test]
205 fn silent_driver_contributes_nothing() {
206 assert_eq!(modulator_signal(30.0, false, false, 10.0, 10.0), 0.0);
207 }
208
209 #[test]
210 fn ungraded_firing_uses_full_magnitude() {
211 assert!((modulator_signal(30.0, true, false, 1.0, 10.0) - 0.3).abs() < 1e-6);
212 assert!((modulator_signal(-80.0, true, false, 1.0, 10.0) + 0.8).abs() < 1e-6);
213 }
214
215 #[test]
216 fn graded_firing_scales_with_membrane_potential_and_caps_at_one() {
217 let half = modulator_signal(30.0, true, true, 5.0, 10.0);
218 assert!((half - 0.15).abs() < 1e-6);
219 let capped = modulator_signal(30.0, true, true, 40.0, 10.0);
220 assert!((capped - 0.3).abs() < 1e-6);
221 }
222
223 #[test]
224 fn factors_multiply_and_reward_adds() {
225 let factor = multiplicative_factor(&[0.3, -0.5]);
226 assert!((factor - (1.3 * 0.5)).abs() < 1e-6);
227 assert!((summed_reward(&[0.3, -0.5]) + 0.2).abs() < 1e-6);
228 }
229
230 #[test]
231 fn leak_and_excitability_clamp_to_unit_interval() {
232 assert!((scale_baseline(ModulatorKind::Leak, 0.8, 2.0) - 1.0).abs() < 1e-6);
233 assert!((scale_baseline(ModulatorKind::FiringProbability, 0.2, 0.0) - 0.0).abs() < 1e-6);
234 assert!(scale_baseline(ModulatorKind::FiringThreshold, 2.0, 3.0) > 0.0);
235 }
236
237 #[test]
238 fn magnitude_below_minus_100_is_rejected() {
239 assert!(validate_instance_fields(-100.1, 1, false, None).is_err());
240 assert!(validate_instance_fields(-100.0, 1, false, None).is_ok());
241 }
242
243 #[test]
244 fn graded_requires_positive_full_scale() {
245 assert!(validate_instance_fields(10.0, 1, true, None).is_err());
246 assert!(validate_instance_fields(10.0, 1, true, Some(0.0)).is_err());
247 assert!(validate_instance_fields(10.0, 1, true, Some(4.0)).is_ok());
248 }
249
250 #[test]
251 fn spike_train_rejects_unlimited_consecutive_fire() {
252 assert!(validate_spike_train(true, 0).is_err());
253 assert!(validate_spike_train(true, 1).is_ok());
254 assert!(validate_spike_train(false, 0).is_ok());
255 }
256
257 #[test]
258 fn type_strings_round_trip() {
259 for kind in ModulatorKind::all() {
260 assert_eq!(ModulatorKind::parse(kind.as_str()).unwrap(), *kind);
261 }
262 }
263}