feagi_structures/neuron_voxels/
class_potential.rs1use crate::FeagiDataError;
11
12pub const MAX_CLASS_COUNT: u32 = 9_999;
15
16pub fn encode_class_potential(class_id: u32, class_count: u32) -> Result<f32, FeagiDataError> {
18 validate_class_count(class_count)?;
19 if class_id >= class_count {
20 return Err(FeagiDataError::BadParameters(format!(
21 "class id {class_id} is outside class count {class_count}"
22 )));
23 }
24 Ok((class_id + 1) as f32 / class_count as f32)
25}
26
27pub fn decode_class_potential(potential: f32, class_count: u32) -> Option<u32> {
31 if class_count == 0 || class_count > MAX_CLASS_COUNT || !potential.is_finite() {
32 return None;
33 }
34 let level = (potential * class_count as f32).round();
35 if level < 1.0 || level > class_count as f32 {
36 return None;
37 }
38 Some(level as u32 - 1)
39}
40
41pub fn validate_class_count(class_count: u32) -> Result<(), FeagiDataError> {
43 if class_count == 0 || class_count > MAX_CLASS_COUNT {
44 return Err(FeagiDataError::BadParameters(format!(
45 "class count must be in 1..={MAX_CLASS_COUNT}, got {class_count}"
46 )));
47 }
48 Ok(())
49}
50
51#[cfg(test)]
52mod tests {
53 use super::*;
54
55 #[test]
56 fn every_class_round_trips() {
57 for class_count in [1_u32, 2, 19, 256, MAX_CLASS_COUNT] {
58 for class_id in [0, class_count / 2, class_count - 1] {
59 let potential = encode_class_potential(class_id, class_count).unwrap();
60 assert!(potential > 0.0 && potential <= 1.0);
61 assert_eq!(
62 decode_class_potential(potential, class_count),
63 Some(class_id)
64 );
65 }
66 }
67 }
68
69 #[test]
70 fn smallest_class_survives_misc_cutoff() {
71 let potential = encode_class_potential(0, MAX_CLASS_COUNT).unwrap();
72 assert!(potential > 1.0e-4);
73 }
74
75 #[test]
76 fn zero_and_out_of_range_decode_to_none() {
77 assert_eq!(decode_class_potential(0.0, 19), None);
78 assert_eq!(decode_class_potential(1.2, 19), None);
79 assert_eq!(decode_class_potential(-0.5, 19), None);
80 assert_eq!(decode_class_potential(f32::NAN, 19), None);
81 assert_eq!(decode_class_potential(0.5, 0), None);
82 }
83
84 #[test]
85 fn small_drift_still_decodes() {
86 let potential = encode_class_potential(13, 19).unwrap();
87 assert_eq!(decode_class_potential(potential + 0.01, 19), Some(13));
88 assert_eq!(decode_class_potential(potential - 0.01, 19), Some(13));
89 }
90
91 #[test]
92 fn rejects_bad_inputs() {
93 assert!(encode_class_potential(19, 19).is_err());
94 assert!(encode_class_potential(0, 0).is_err());
95 assert!(encode_class_potential(0, MAX_CLASS_COUNT + 1).is_err());
96 }
97}