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