Skip to main content

feagi_structures/genomic/classifiers/
mod.rs

1// Copyright 2025 Neuraville Inc.
2// SPDX-License-Identifier: Apache-2.0
3
4/*!
5First-class genome classifier assembly.
6
7Classifiers are stored under the top-level genome key `classifiers`, parallel
8to `brain_regions`. A classifier is not a brain region and is not exportable
9as a circuit. It records the assembly's properties, owned internals, and
10referenced input areas so neuroembryogenesis and area/mapping edits stay
11aligned.
12
13Each field binding is one Classifier mapping: an interconnect area scanning
14the shared kernel memory, with its own detection twin.
15*/
16
17use crate::neuron_voxels::class_potential::validate_class_count;
18use serde::{Deserialize, Serialize};
19use std::collections::HashMap;
20
21/// Morphology used for kernel → kernel-memory encode.
22pub const CLASSIFIER_KERNEL_MORPHOLOGY: &str = "episodic_memory";
23/// Morphology used for class → class-memory encode.
24pub const CLASSIFIER_CLASS_MORPHOLOGY: &str = "episodic_memory";
25/// Morphology used for kernel-memory → class-memory bind.
26pub const CLASSIFIER_ASSOCIATIVE_MORPHOLOGY: &str = "associative_memory";
27/// Morphology used for field → kernel-memory scan.
28pub const CLASSIFIER_SCAN_MORPHOLOGY: &str = "episodic_scan";
29
30/// Directed mapping owned by a classifier assembly.
31#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct ClassifierMapping {
33    pub src_area_id: String,
34    pub dst_area_id: String,
35    pub morphology_id: String,
36}
37
38/// One field area scanning this classifier, and the twin that shows its detections.
39#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
40pub struct ClassifierField {
41    pub field_area_id: String,
42    pub scan_twin_id: String,
43}
44
45/// How a classifier learns. Recall uses the same kernel geometry as training.
46#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)]
47#[serde(rename_all = "snake_case")]
48pub enum ClassifierTrainingMode {
49    /// One kernel sample and one class sample per burst.
50    #[default]
51    Kernel,
52    /// Slide `kernel_size` across each mapped field and label it from the mask.
53    Scanner,
54}
55
56impl ClassifierTrainingMode {
57    /// Genome and API spelling.
58    pub fn as_str(self) -> &'static str {
59        match self {
60            Self::Kernel => "kernel",
61            Self::Scanner => "scanner",
62        }
63    }
64
65    /// Parse a stored mode. Absent values are not accepted here.
66    pub fn parse(value: &str) -> Result<Self, String> {
67        match value {
68            "kernel" => Ok(Self::Kernel),
69            "scanner" => Ok(Self::Scanner),
70            other => Err(format!(
71                "training_mode must be \"kernel\" or \"scanner\", got \"{other}\""
72            )),
73        }
74    }
75}
76
77/// First-class classifier record persisted in the genome.
78#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
79pub struct Classifier {
80    /// UUID string, genome map key.
81    pub classifier_id: String,
82    pub name: String,
83    /// Region that contains this assembly. Classifiers are not regions.
84    pub parent_region_id: String,
85    pub coordinates_3d: [i32; 3],
86    /// Kernel mode ingests one sample per burst. Scanner mode learns from field windows.
87    /// Genomes saved before modes existed load as kernel mode.
88    #[serde(default)]
89    pub training_mode: ClassifierTrainingMode,
90    /// Referenced inputs. Cleared when that area is deleted.
91    #[serde(default, skip_serializing_if = "Option::is_none")]
92    pub kernel_area_id: Option<String>,
93    /// Kernel-mode class input. Must be `1×1×n`: each depth voxel is one class channel.
94    #[serde(default, skip_serializing_if = "Option::is_none")]
95    pub class_area_id: Option<String>,
96    /// Scanner-mode label plane, `W×H×1`. Width and height match each mapped field.
97    /// Each pixel's potential is `(class_id + 1) / class_count`.
98    #[serde(default, skip_serializing_if = "Option::is_none")]
99    pub mask_area_id: Option<String>,
100    /// Scanner-mode class count. Decodes mask potentials and encodes detection twins.
101    #[serde(default, skip_serializing_if = "Option::is_none")]
102    pub class_count: Option<u32>,
103    /// Scanner-mode kernel `[x, y, z]`. Z must equal each mapped field's depth.
104    #[serde(default, skip_serializing_if = "Option::is_none")]
105    pub kernel_size: Option<[u32; 3]>,
106    /// Field scans. Empty until Classifier mappings are drawn.
107    #[serde(default)]
108    pub fields: Vec<ClassifierField>,
109    /// Owned internals. Deleting kernel or class memory deletes the classifier.
110    pub kernel_memory_id: String,
111    pub class_memory_id: String,
112    /// Per-scanning-instance pain and pleasure on the associative mapping.
113    /// Genomes saved before reward training existed stay off.
114    #[serde(default)]
115    pub reward_training: bool,
116    /// Correct-answer area. Scanner mode matches each detection twin.
117    /// Kernel mode matches the class area. Empty until the user selects one.
118    #[serde(default, skip_serializing_if = "Option::is_none")]
119    pub answer_feedback_area_id: Option<String>,
120    /// Hidden area whose firing is this classifier's pain signal.
121    #[serde(default, skip_serializing_if = "Option::is_none")]
122    pub pain_area_id: Option<String>,
123    /// Hidden area whose firing is this classifier's pleasure signal.
124    #[serde(default, skip_serializing_if = "Option::is_none")]
125    pub pleasure_area_id: Option<String>,
126    /// Bursts between a decision and the answer that grades it. Zero grades the same burst.
127    #[serde(default)]
128    pub answer_latency_bursts: u32,
129    /// When set, pain and pleasure run only on bursts this area fires.
130    #[serde(default, skip_serializing_if = "Option::is_none")]
131    pub learn_area_id: Option<String>,
132    /// OPU that receives one surplus value per class channel.
133    #[serde(default, skip_serializing_if = "Option::is_none")]
134    pub confidence_area_id: Option<String>,
135    #[serde(default)]
136    pub properties: HashMap<String, serde_json::Value>,
137}
138
139impl Classifier {
140    /// Kernel memory and class memory. Deleting either deletes the assembly.
141    pub fn assembly_core_ids(&self) -> Vec<String> {
142        vec![self.kernel_memory_id.clone(), self.class_memory_id.clone()]
143    }
144
145    /// Owned internals deleted with the classifier, including every field twin.
146    pub fn owned_area_ids(&self) -> Vec<String> {
147        let mut owned = self.assembly_core_ids();
148        for field in &self.fields {
149            if !field.scan_twin_id.is_empty() {
150                owned.push(field.scan_twin_id.clone());
151            }
152        }
153        if let Some(pain_area_id) = &self.pain_area_id {
154            if !pain_area_id.is_empty() {
155                owned.push(pain_area_id.clone());
156            }
157        }
158        if let Some(pleasure_area_id) = &self.pleasure_area_id {
159            if !pleasure_area_id.is_empty() {
160                owned.push(pleasure_area_id.clone());
161            }
162        }
163        owned
164    }
165
166    /// Referenced input areas that may be cleared independently.
167    pub fn input_area_ids(&self) -> Vec<String> {
168        let mut inputs = Vec::new();
169        if let Some(kernel) = &self.kernel_area_id {
170            inputs.push(kernel.clone());
171        }
172        if let Some(class) = &self.class_area_id {
173            inputs.push(class.clone());
174        }
175        if let Some(mask) = &self.mask_area_id {
176            inputs.push(mask.clone());
177        }
178        for field in &self.fields {
179            inputs.push(field.field_area_id.clone());
180        }
181        inputs
182    }
183
184    pub fn owns_assembly_core(&self, area_id: &str) -> bool {
185        self.kernel_memory_id == area_id || self.class_memory_id == area_id
186    }
187
188    /// True when `area_id` is kernel memory, class memory, or one of the field twins.
189    pub fn owns_area(&self, area_id: &str) -> bool {
190        self.owns_assembly_core(area_id) || self.field_for_twin(area_id).is_some()
191    }
192
193    pub fn field_for_twin(&self, twin_id: &str) -> Option<&ClassifierField> {
194        self.fields
195            .iter()
196            .find(|field| field.scan_twin_id == twin_id)
197    }
198
199    pub fn binding_for_field(&self, field_area_id: &str) -> Option<&ClassifierField> {
200        self.fields
201            .iter()
202            .find(|field| field.field_area_id == field_area_id)
203    }
204
205    /// True when `area_id` is a referenced input of this classifier.
206    pub fn references_input(&self, area_id: &str) -> bool {
207        self.kernel_area_id.as_deref() == Some(area_id)
208            || self.class_area_id.as_deref() == Some(area_id)
209            || self.mask_area_id.as_deref() == Some(area_id)
210            || self.binding_for_field(area_id).is_some()
211    }
212
213    /// Replace the training mode and its inputs. The other mode's slots are cleared.
214    ///
215    /// Returns true when long-term memory learned under the previous geometry must be dropped.
216    pub fn apply_training_inputs(
217        &mut self,
218        mode: ClassifierTrainingMode,
219        kernel_area_id: Option<String>,
220        class_area_id: Option<String>,
221        mask_area_id: Option<String>,
222        kernel_size: Option<[u32; 3]>,
223        class_count: Option<u32>,
224    ) -> Result<bool, String> {
225        let previous_mode = self.training_mode;
226        let previous_size = self.kernel_size;
227        let previous_class_count = self.class_count;
228        match mode {
229            ClassifierTrainingMode::Kernel => {
230                let kernel = kernel_area_id.ok_or_else(|| "kernel_area_id required".to_string())?;
231                let class = class_area_id.ok_or_else(|| "class_area_id required".to_string())?;
232                self.kernel_area_id = Some(required_area_id(kernel, "kernel_area_id")?);
233                self.class_area_id = Some(required_area_id(class, "class_area_id")?);
234                self.mask_area_id = None;
235                self.kernel_size = None;
236                self.class_count = None;
237            }
238            ClassifierTrainingMode::Scanner => {
239                let mask = mask_area_id.ok_or_else(|| "mask_area_id required".to_string())?;
240                let size = kernel_size.ok_or_else(|| "kernel_size required".to_string())?;
241                let count = class_count.ok_or_else(|| "class_count required".to_string())?;
242                validate_kernel_size(size)?;
243                validate_class_count(count).map_err(|e| e.to_string())?;
244                self.mask_area_id = Some(required_area_id(mask, "mask_area_id")?);
245                self.kernel_size = Some(size);
246                self.class_count = Some(count);
247                self.kernel_area_id = None;
248                self.class_area_id = None;
249            }
250        }
251        self.training_mode = mode;
252        Ok(previous_mode != mode
253            || (mode == ClassifierTrainingMode::Scanner
254                && (previous_size != self.kernel_size || previous_class_count != self.class_count)))
255    }
256
257    /// Apply classifier-level edit. Does not replace owned internals or field bindings.
258    pub fn apply_assembly_update(
259        &mut self,
260        name: Option<String>,
261        coordinates_3d: Option<[i32; 3]>,
262        parent_region_id: Option<String>,
263        kernel_area_id: Option<String>,
264        class_area_id: Option<String>,
265    ) -> Result<(), String> {
266        if let Some(name) = name {
267            let trimmed = name.trim();
268            if trimmed.is_empty() {
269                return Err("Classifier name cannot be blank".to_string());
270            }
271            self.name = trimmed.to_string();
272        }
273        if let Some(coordinates_3d) = coordinates_3d {
274            self.coordinates_3d = coordinates_3d;
275        }
276        if let Some(parent_region_id) = parent_region_id {
277            let trimmed = parent_region_id.trim();
278            if trimmed.is_empty() {
279                return Err("parent_region_id cannot be blank".to_string());
280            }
281            self.parent_region_id = trimmed.to_string();
282        }
283        if let Some(kernel_area_id) = kernel_area_id {
284            self.kernel_area_id = Some(required_area_id(kernel_area_id, "kernel_area_id")?);
285        }
286        if let Some(class_area_id) = class_area_id {
287            self.class_area_id = Some(required_area_id(class_area_id, "class_area_id")?);
288        }
289        Ok(())
290    }
291
292    /// Apply classifier-level metadata. Does not change owned internals or input slots.
293    pub fn apply_metadata_update(
294        &mut self,
295        name: Option<String>,
296        coordinates_3d: Option<[i32; 3]>,
297    ) -> Result<(), String> {
298        self.apply_assembly_update(name, coordinates_3d, None, None, None)
299    }
300
301    /// Mappings this classifier requires given its current inputs.
302    pub fn required_mappings(&self) -> Vec<ClassifierMapping> {
303        let mut mappings = Vec::new();
304        if let Some(kernel) = &self.kernel_area_id {
305            mappings.push(ClassifierMapping {
306                src_area_id: kernel.clone(),
307                dst_area_id: self.kernel_memory_id.clone(),
308                morphology_id: CLASSIFIER_KERNEL_MORPHOLOGY.to_string(),
309            });
310        }
311        if let Some(class) = &self.class_area_id {
312            mappings.push(ClassifierMapping {
313                src_area_id: class.clone(),
314                dst_area_id: self.class_memory_id.clone(),
315                morphology_id: CLASSIFIER_CLASS_MORPHOLOGY.to_string(),
316            });
317        }
318        mappings.push(ClassifierMapping {
319            src_area_id: self.kernel_memory_id.clone(),
320            dst_area_id: self.class_memory_id.clone(),
321            morphology_id: CLASSIFIER_ASSOCIATIVE_MORPHOLOGY.to_string(),
322        });
323        for field in &self.fields {
324            mappings.push(ClassifierMapping {
325                src_area_id: field.field_area_id.clone(),
326                dst_area_id: self.kernel_memory_id.clone(),
327                morphology_id: CLASSIFIER_SCAN_MORPHOLOGY.to_string(),
328            });
329        }
330        mappings
331    }
332
333    /// Drop an input reference when that area is deleted.
334    pub fn clear_input(&mut self, area_id: &str) {
335        if self.kernel_area_id.as_deref() == Some(area_id) {
336            self.kernel_area_id = None;
337        }
338        if self.class_area_id.as_deref() == Some(area_id) {
339            self.class_area_id = None;
340        }
341        if self.mask_area_id.as_deref() == Some(area_id) {
342            self.mask_area_id = None;
343        }
344        if self.answer_feedback_area_id.as_deref() == Some(area_id) {
345            self.answer_feedback_area_id = None;
346        }
347        if self.learn_area_id.as_deref() == Some(area_id) {
348            self.learn_area_id = None;
349        }
350        if self.confidence_area_id.as_deref() == Some(area_id) {
351            self.confidence_area_id = None;
352        }
353        self.fields.retain(|field| field.field_area_id != area_id);
354    }
355
356    /// Remove one field binding. Returns the removed twin id.
357    pub fn detach_field(&mut self, field_area_id: &str) -> Option<String> {
358        let position = self
359            .fields
360            .iter()
361            .position(|field| field.field_area_id == field_area_id)?;
362        Some(self.fields.remove(position).scan_twin_id)
363    }
364
365    /// Remove the binding whose twin is `twin_id`. Returns the field area id.
366    pub fn detach_twin(&mut self, twin_id: &str) -> Option<String> {
367        let position = self
368            .fields
369            .iter()
370            .position(|field| field.scan_twin_id == twin_id)?;
371        Some(self.fields.remove(position).field_area_id)
372    }
373
374    pub fn attach_field(
375        &mut self,
376        field_area_id: String,
377        scan_twin_id: String,
378    ) -> Result<(), String> {
379        let field_area_id = required_area_id(field_area_id, "field_area_id")?;
380        let scan_twin_id = required_area_id(scan_twin_id, "scan_twin_id")?;
381        if self.binding_for_field(&field_area_id).is_some() {
382            return Err(format!(
383                "field_area_id {field_area_id} is already mapped to this classifier"
384            ));
385        }
386        self.fields.push(ClassifierField {
387            field_area_id,
388            scan_twin_id,
389        });
390        Ok(())
391    }
392
393    /// Bind or clear kernel/class inputs from a mapping change. Field scans are attached explicitly.
394    pub fn apply_mapping_change(
395        &mut self,
396        src_area_id: &str,
397        dst_area_id: &str,
398        morphology_id: &str,
399        removed: bool,
400    ) -> bool {
401        if dst_area_id == self.kernel_memory_id && morphology_id == CLASSIFIER_KERNEL_MORPHOLOGY {
402            self.kernel_area_id = if removed {
403                None
404            } else {
405                Some(src_area_id.to_string())
406            };
407            return true;
408        }
409        if dst_area_id == self.class_memory_id && morphology_id == CLASSIFIER_CLASS_MORPHOLOGY {
410            self.class_area_id = if removed {
411                None
412            } else {
413                Some(src_area_id.to_string())
414            };
415            return true;
416        }
417        false
418    }
419}
420
421/// Kernel-mode class input. Width and height are 1 so depth index `z` is twin class `z`.
422pub fn validate_kernel_class_area_shape(dimensions: [u32; 3]) -> Result<(), String> {
423    if dimensions[0] != 1 || dimensions[1] != 1 || dimensions[2] == 0 {
424        return Err(
425            "class area must be 1x1xn so each depth voxel maps to one detection class".to_string(),
426        );
427    }
428    Ok(())
429}
430
431/// Every kernel axis must be at least one voxel.
432pub fn validate_kernel_size(size: [u32; 3]) -> Result<(), String> {
433    if size[0] == 0 || size[1] == 0 || size[2] == 0 {
434        return Err("kernel_size axes must be greater than zero".to_string());
435    }
436    Ok(())
437}
438
439/// Every detection twin is one layer over its field: `field_w × field_h × 1`.
440pub fn detection_twin_shape(field: [u32; 3]) -> [u32; 3] {
441    [field[0], field[1], 1]
442}
443
444/// Scanner kernel Z matches the image depth. The mask is one layer with the image's width and height.
445pub fn validate_scanner_field(
446    kernel_size: [u32; 3],
447    field: [u32; 3],
448    mask: [u32; 3],
449) -> Result<(), String> {
450    validate_kernel_size(kernel_size)?;
451    if field[0] == 0 || field[1] == 0 || field[2] == 0 {
452        return Err("field dimensions must be greater than zero".to_string());
453    }
454    if mask[2] != 1 {
455        return Err("mask must be one layer deep; class ids are carried as potential".to_string());
456    }
457    if kernel_size[0] > field[0] || kernel_size[1] > field[1] {
458        return Err("kernel_size does not fit the field".to_string());
459    }
460    if kernel_size[2] != field[2] {
461        return Err("kernel depth must equal the field depth".to_string());
462    }
463    if mask[0] != field[0] || mask[1] != field[1] {
464        return Err("mask width and height must equal the field".to_string());
465    }
466    Ok(())
467}
468
469/// One class for the whole image: two axes are 1 and the remaining axis is the class count.
470pub fn is_whole_image_class_shape(feedback: [u32; 3], class_count: u32) -> bool {
471    if class_count == 0 {
472        return false;
473    }
474    let ones = feedback.iter().filter(|axis| **axis == 1).count();
475    let volume = u64::from(feedback[0])
476        .saturating_mul(u64::from(feedback[1]))
477        .saturating_mul(u64::from(feedback[2]));
478    ones >= 2 && volume == u64::from(class_count)
479}
480
481/// Answer feedback must match the classifier output.
482///
483/// Kernel mode compares the class area. Scanner mode compares each detection
484/// twin (`W×H×1`, class as potential) or accepts a whole-image class of
485/// `class_count` channels. With no twin yet, the mask is that output.
486pub fn validate_answer_feedback_shape(
487    mode: ClassifierTrainingMode,
488    feedback: [u32; 3],
489    reference: [u32; 3],
490    output_shapes: &[[u32; 3]],
491    class_count: u32,
492) -> Result<(), String> {
493    if feedback[0] == 0 || feedback[1] == 0 || feedback[2] == 0 {
494        return Err("answer feedback dimensions must be greater than zero".to_string());
495    }
496    match mode {
497        ClassifierTrainingMode::Kernel => {
498            if feedback != reference {
499                return Err("answer feedback dimensions must match the class area".to_string());
500            }
501        }
502        ClassifierTrainingMode::Scanner => {
503            if is_whole_image_class_shape(feedback, class_count) {
504                return Ok(());
505            }
506            let shapes = if output_shapes.is_empty() {
507                std::slice::from_ref(&reference)
508            } else {
509                output_shapes
510            };
511            if shapes.iter().any(|shape| *shape != feedback) {
512                return Err(
513                    "answer feedback dimensions must match the detection output or a whole-image class"
514                        .to_string(),
515                );
516            }
517        }
518    }
519    Ok(())
520}
521
522fn required_area_id(area_id: String, field: &str) -> Result<String, String> {
523    let trimmed = area_id.trim();
524    if trimmed.is_empty() {
525        return Err(format!("{field} cannot be blank"));
526    }
527    Ok(trimmed.to_string())
528}
529
530#[cfg(test)]
531mod tests {
532    use super::*;
533
534    fn sample() -> Classifier {
535        let mut classifier = Classifier {
536            classifier_id: "clf-1".to_string(),
537            name: "demo".to_string(),
538            parent_region_id: "region".to_string(),
539            coordinates_3d: [1, 2, 3],
540            training_mode: ClassifierTrainingMode::Kernel,
541            kernel_area_id: Some("kernel".to_string()),
542            class_area_id: Some("class".to_string()),
543            mask_area_id: None,
544            class_count: None,
545            kernel_size: None,
546            fields: Vec::new(),
547            kernel_memory_id: "kmem".to_string(),
548            class_memory_id: "cmem".to_string(),
549            reward_training: false,
550            answer_feedback_area_id: None,
551            pain_area_id: None,
552            pleasure_area_id: None,
553            answer_latency_bursts: 0,
554            learn_area_id: None,
555            confidence_area_id: None,
556            properties: HashMap::new(),
557        };
558        classifier
559            .attach_field("field".to_string(), "twin".to_string())
560            .expect("first field");
561        classifier
562    }
563
564    #[test]
565    fn required_mappings_cover_shared_edges_and_each_field() {
566        let mut classifier = sample();
567        classifier
568            .attach_field("field-b".to_string(), "twin-b".to_string())
569            .expect("second field");
570        let mappings = classifier.required_mappings();
571        assert_eq!(mappings.len(), 5);
572        assert!(mappings.iter().any(|m| {
573            m.src_area_id == "field"
574                && m.dst_area_id == "kmem"
575                && m.morphology_id == CLASSIFIER_SCAN_MORPHOLOGY
576        }));
577        assert!(mappings.iter().any(|m| m.src_area_id == "field-b"));
578    }
579
580    #[test]
581    fn deleting_one_field_keeps_the_other_eye() {
582        let mut classifier = sample();
583        classifier
584            .attach_field("field-b".to_string(), "twin-b".to_string())
585            .expect("second field");
586        assert_eq!(classifier.detach_field("field"), Some("twin".to_string()));
587        assert!(classifier.binding_for_field("field").is_none());
588        assert_eq!(
589            classifier
590                .binding_for_field("field-b")
591                .map(|f| f.scan_twin_id.as_str()),
592            Some("twin-b")
593        );
594        assert!(classifier.owns_assembly_core("kmem"));
595        assert!(!classifier.owns_area("twin"));
596        assert!(classifier.owns_area("twin-b"));
597    }
598
599    #[test]
600    fn duplicate_field_mapping_is_rejected() {
601        let mut classifier = sample();
602        let result = classifier.attach_field("field".to_string(), "other-twin".to_string());
603        assert!(result.is_err());
604        assert_eq!(classifier.fields.len(), 1);
605    }
606
607    #[test]
608    fn metadata_update_renames_without_touching_areas() {
609        let mut classifier = sample();
610        classifier
611            .apply_metadata_update(Some("  renamed  ".to_string()), Some([9, 8, 7]))
612            .expect("valid metadata");
613        assert_eq!(classifier.name, "renamed");
614        assert_eq!(classifier.coordinates_3d, [9, 8, 7]);
615        assert_eq!(classifier.kernel_memory_id, "kmem");
616        assert_eq!(classifier.fields[0].scan_twin_id, "twin");
617        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
618    }
619
620    #[test]
621    fn assembly_update_retargets_kernel_and_class_only() {
622        let mut classifier = sample();
623        classifier
624            .apply_assembly_update(
625                None,
626                None,
627                Some("other-region".to_string()),
628                Some("kernel2".to_string()),
629                Some("class2".to_string()),
630            )
631            .expect("valid assembly update");
632        assert_eq!(classifier.parent_region_id, "other-region");
633        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel2"));
634        assert_eq!(classifier.class_area_id.as_deref(), Some("class2"));
635        assert_eq!(classifier.fields[0].field_area_id, "field");
636    }
637
638    #[test]
639    fn scanner_mode_clears_kernel_inputs_and_reports_geometry_change() {
640        let mut classifier = sample();
641        let changed = classifier
642            .apply_training_inputs(
643                ClassifierTrainingMode::Scanner,
644                None,
645                None,
646                Some("mask".to_string()),
647                Some([8, 8, 3]),
648                Some(19),
649            )
650            .expect("scanner inputs");
651        assert!(changed);
652        assert_eq!(classifier.training_mode, ClassifierTrainingMode::Scanner);
653        assert!(classifier.kernel_area_id.is_none());
654        assert!(classifier.class_area_id.is_none());
655        assert_eq!(classifier.mask_area_id.as_deref(), Some("mask"));
656        assert_eq!(classifier.kernel_size, Some([8, 8, 3]));
657        assert_eq!(classifier.class_count, Some(19));
658        assert!(classifier.references_input("mask"));
659        let same = classifier
660            .apply_training_inputs(
661                ClassifierTrainingMode::Scanner,
662                None,
663                None,
664                Some("mask".to_string()),
665                Some([8, 8, 3]),
666                Some(19),
667            )
668            .expect("same scanner geometry");
669        assert!(!same);
670        let recounted = classifier
671            .apply_training_inputs(
672                ClassifierTrainingMode::Scanner,
673                None,
674                None,
675                Some("mask".to_string()),
676                Some([8, 8, 3]),
677                Some(6),
678            )
679            .expect("new class count");
680        assert!(recounted, "a new class count changes what learned ids mean");
681    }
682
683    #[test]
684    fn scanner_mode_requires_a_valid_class_count() {
685        let mut classifier = sample();
686        for bad in [None, Some(0)] {
687            assert!(classifier
688                .apply_training_inputs(
689                    ClassifierTrainingMode::Scanner,
690                    None,
691                    None,
692                    Some("mask".to_string()),
693                    Some([8, 8, 3]),
694                    bad,
695                )
696                .is_err());
697        }
698        assert_eq!(classifier.training_mode, ClassifierTrainingMode::Kernel);
699    }
700
701    #[test]
702    fn kernel_mode_clears_scanner_inputs() {
703        let mut classifier = sample();
704        classifier
705            .apply_training_inputs(
706                ClassifierTrainingMode::Scanner,
707                None,
708                None,
709                Some("mask".to_string()),
710                Some([2, 2, 1]),
711                Some(4),
712            )
713            .expect("scanner");
714        classifier
715            .apply_training_inputs(
716                ClassifierTrainingMode::Kernel,
717                Some("kernel".to_string()),
718                Some("class".to_string()),
719                None,
720                None,
721                None,
722            )
723            .expect("kernel");
724        assert!(classifier.mask_area_id.is_none());
725        assert!(classifier.kernel_size.is_none());
726        assert!(classifier.class_count.is_none());
727        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
728    }
729
730    #[test]
731    fn detection_twin_is_one_layer_over_the_field() {
732        assert_eq!(detection_twin_shape([256, 128, 3]), [256, 128, 1]);
733    }
734
735    #[test]
736    fn answer_feedback_matches_class_area_or_detection_output() {
737        assert!(validate_answer_feedback_shape(
738            ClassifierTrainingMode::Kernel,
739            [1, 1, 4],
740            [1, 1, 4],
741            &[],
742            4,
743        )
744        .is_ok());
745        assert!(validate_answer_feedback_shape(
746            ClassifierTrainingMode::Kernel,
747            [8, 8, 4],
748            [1, 1, 4],
749            &[],
750            4,
751        )
752        .is_err());
753        assert!(validate_answer_feedback_shape(
754            ClassifierTrainingMode::Scanner,
755            [16, 16, 1],
756            [16, 16, 1],
757            &[[16, 16, 1]],
758            4,
759        )
760        .is_ok());
761        assert!(validate_answer_feedback_shape(
762            ClassifierTrainingMode::Scanner,
763            [16, 16, 1],
764            [16, 16, 1],
765            &[[16, 16, 1], [8, 8, 1]],
766            4,
767        )
768        .is_err());
769        assert!(
770            validate_answer_feedback_shape(
771                ClassifierTrainingMode::Scanner,
772                [16, 16, 4],
773                [16, 16, 1],
774                &[[16, 16, 1]],
775                4,
776            )
777            .is_err(),
778            "a one-hot class volume no longer matches a potential-coded twin"
779        );
780        assert!(validate_answer_feedback_shape(
781            ClassifierTrainingMode::Scanner,
782            [10, 1, 1],
783            [16, 16, 1],
784            &[[16, 16, 1]],
785            10,
786        )
787        .is_ok());
788    }
789
790    #[test]
791    fn kernel_class_area_must_be_one_by_one_by_n() {
792        assert!(validate_kernel_class_area_shape([1, 1, 4]).is_ok());
793        assert!(validate_kernel_class_area_shape([2, 7, 4]).is_err());
794        assert!(validate_kernel_class_area_shape([4, 1, 1]).is_err());
795        assert!(validate_kernel_class_area_shape([1, 1, 0]).is_err());
796    }
797
798    #[test]
799    fn scanner_field_must_match_mask_and_kernel_depth() {
800        assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 1]).is_ok());
801        assert!(validate_scanner_field([8, 8, 1], [256, 128, 3], [256, 128, 1]).is_err());
802        assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [200, 128, 1]).is_err());
803        assert!(
804            validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 10]).is_err(),
805            "one-hot class masks are replaced by the single-layer potential mask"
806        );
807    }
808}