Skip to main content

feagi_structures/neuron_voxels/
class_potential.rs

1//! Class-id encoding for single-layer (`W×H×1`) class maps.
2//!
3//! Object segmentation input, object segmentation output, and classifier
4//! detection twins carry one voxel per pixel. The class is the potential:
5//! `(class_id + 1) / class_count`, always in `(0, 1]`. Zero means unlabeled,
6//! so the Misc codec's `[-1, 1]` range and near-zero cutoff both hold.
7//!
8//! @cursor:ffi-safe - pure arithmetic, no allocation.
9
10use crate::FeagiDataError;
11
12/// Largest class count the encoding supports. Misc drops `|p| <= 1e-4`, so the
13/// smallest encoded class `1 / class_count` must stay above that cutoff.
14pub const MAX_CLASS_COUNT: u32 = 9_999;
15
16/// Potential for `class_id` in a map of `class_count` classes.
17pub 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
27/// Class id carried by `potential`, or `None` when it names no class.
28///
29/// Rounds to the nearest class so values that picked up float error still decode.
30pub 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
41/// Class counts must be in `1..=MAX_CLASS_COUNT`.
42pub 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}