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 class output.
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/// Scanner class output: one layer over the field, class carried as potential.
440pub fn detection_twin_shape(field: [u32; 3]) -> [u32; 3] {
441    [field[0], field[1], 1]
442}
443
444/// Kernel mode compares the whole field to the kernel. The field must be that size.
445pub fn validate_kernel_field(kernel: [u32; 3], field: [u32; 3]) -> Result<(), String> {
446    if kernel[0] == 0 || kernel[1] == 0 || kernel[2] == 0 {
447        return Err("kernel area dimensions must be greater than zero".to_string());
448    }
449    if kernel != field {
450        return Err(format!(
451            "kernel mode field must match the kernel area ({}x{}x{}); got {}x{}x{}",
452            kernel[0], kernel[1], kernel[2], field[0], field[1], field[2]
453        ));
454    }
455    Ok(())
456}
457
458/// Class output geometry for one training mode.
459///
460/// Kernel mode is `1×1×class_count`, the same shape as the class input.
461/// Scanner mode is `field_w × field_h × 1`.
462pub fn class_output_shape(
463    mode: ClassifierTrainingMode,
464    field: [u32; 3],
465    class_count: u32,
466) -> Result<[u32; 3], String> {
467    match mode {
468        ClassifierTrainingMode::Kernel => {
469            validate_class_count(class_count).map_err(|error| error.to_string())?;
470            Ok([1, 1, class_count])
471        }
472        ClassifierTrainingMode::Scanner => {
473            if field[0] == 0 || field[1] == 0 {
474                return Err("scanner class output needs a field width and height".to_string());
475            }
476            Ok(detection_twin_shape(field))
477        }
478    }
479}
480
481/// Scanner outputs forward the class potential. Kernel outputs fire depth `z`.
482pub fn class_output_forwards_potential(mode: ClassifierTrainingMode) -> bool {
483    matches!(mode, ClassifierTrainingMode::Scanner)
484}
485
486/// Visible name of one class output. A second field includes that field's name.
487pub fn class_output_area_name(
488    classifier_name: &str,
489    field_name: &str,
490    output_count: usize,
491) -> String {
492    if output_count <= 1 {
493        format!("{classifier_name} class output")
494    } else {
495        format!("{classifier_name} {field_name} class output")
496    }
497}
498
499/// Scanner kernel Z matches the image depth. The mask is one layer with the image's width and height.
500pub fn validate_scanner_field(
501    kernel_size: [u32; 3],
502    field: [u32; 3],
503    mask: [u32; 3],
504) -> Result<(), String> {
505    validate_kernel_size(kernel_size)?;
506    if field[0] == 0 || field[1] == 0 || field[2] == 0 {
507        return Err("field dimensions must be greater than zero".to_string());
508    }
509    if mask[2] != 1 {
510        return Err("mask must be one layer deep; class ids are carried as potential".to_string());
511    }
512    if kernel_size[0] > field[0] || kernel_size[1] > field[1] {
513        return Err("kernel_size does not fit the field".to_string());
514    }
515    if kernel_size[2] != field[2] {
516        return Err("kernel depth must equal the field depth".to_string());
517    }
518    if mask[0] != field[0] || mask[1] != field[1] {
519        return Err("mask width and height must equal the field".to_string());
520    }
521    Ok(())
522}
523
524/// One class for the whole image: two axes are 1 and the remaining axis is the class count.
525pub fn is_whole_image_class_shape(feedback: [u32; 3], class_count: u32) -> bool {
526    if class_count == 0 {
527        return false;
528    }
529    let ones = feedback.iter().filter(|axis| **axis == 1).count();
530    let volume = u64::from(feedback[0])
531        .saturating_mul(u64::from(feedback[1]))
532        .saturating_mul(u64::from(feedback[2]));
533    ones >= 2 && volume == u64::from(class_count)
534}
535
536/// Answer feedback must match the classifier output.
537///
538/// Kernel mode compares the class area. Scanner mode compares each detection
539/// twin (`W×H×1`, class as potential) or accepts a whole-image class of
540/// `class_count` channels. With no twin yet, the mask is that output.
541pub fn validate_answer_feedback_shape(
542    mode: ClassifierTrainingMode,
543    feedback: [u32; 3],
544    reference: [u32; 3],
545    output_shapes: &[[u32; 3]],
546    class_count: u32,
547) -> Result<(), String> {
548    if feedback[0] == 0 || feedback[1] == 0 || feedback[2] == 0 {
549        return Err("answer feedback dimensions must be greater than zero".to_string());
550    }
551    match mode {
552        ClassifierTrainingMode::Kernel => {
553            if feedback != reference {
554                return Err("answer feedback dimensions must match the class area".to_string());
555            }
556        }
557        ClassifierTrainingMode::Scanner => {
558            if is_whole_image_class_shape(feedback, class_count) {
559                return Ok(());
560            }
561            let shapes = if output_shapes.is_empty() {
562                std::slice::from_ref(&reference)
563            } else {
564                output_shapes
565            };
566            if shapes.iter().any(|shape| *shape != feedback) {
567                return Err(
568                    "answer feedback dimensions must match the detection output or a whole-image class"
569                        .to_string(),
570                );
571            }
572        }
573    }
574    Ok(())
575}
576
577/// Mapping rule written for one classifier-owned edge.
578///
579/// Only `associative_memory` is plastic; its window is the kernel memory's temporal depth.
580pub fn classifier_mapping_rule(morphology_id: &str, associative_window: u32) -> serde_json::Value {
581    let is_associative = morphology_id == CLASSIFIER_ASSOCIATIVE_MORPHOLOGY;
582    let plasticity_value = if is_associative { 1 } else { 0 };
583    serde_json::json!({
584        "morphology_id": morphology_id,
585        "morphology_scalar": [1, 1, 1],
586        "postSynapticCurrent_multiplier": 1,
587        "plasticity_flag": is_associative,
588        "plasticity_constant": plasticity_value,
589        "ltp_multiplier": plasticity_value,
590        "ltd_multiplier": plasticity_value,
591        "plasticity_window": if is_associative { associative_window } else { 0 },
592    })
593}
594
595fn required_area_id(area_id: String, field: &str) -> Result<String, String> {
596    let trimmed = area_id.trim();
597    if trimmed.is_empty() {
598        return Err(format!("{field} cannot be blank"));
599    }
600    Ok(trimmed.to_string())
601}
602
603#[cfg(test)]
604mod tests {
605    use super::*;
606
607    fn sample() -> Classifier {
608        let mut classifier = Classifier {
609            classifier_id: "clf-1".to_string(),
610            name: "demo".to_string(),
611            parent_region_id: "region".to_string(),
612            coordinates_3d: [1, 2, 3],
613            training_mode: ClassifierTrainingMode::Kernel,
614            kernel_area_id: Some("kernel".to_string()),
615            class_area_id: Some("class".to_string()),
616            mask_area_id: None,
617            class_count: None,
618            kernel_size: None,
619            fields: Vec::new(),
620            kernel_memory_id: "kmem".to_string(),
621            class_memory_id: "cmem".to_string(),
622            reward_training: false,
623            answer_feedback_area_id: None,
624            pain_area_id: None,
625            pleasure_area_id: None,
626            answer_latency_bursts: 0,
627            learn_area_id: None,
628            confidence_area_id: None,
629            properties: HashMap::new(),
630        };
631        classifier
632            .attach_field("field".to_string(), "twin".to_string())
633            .expect("first field");
634        classifier
635    }
636
637    #[test]
638    fn required_mappings_cover_shared_edges_and_each_field() {
639        let mut classifier = sample();
640        classifier
641            .attach_field("field-b".to_string(), "twin-b".to_string())
642            .expect("second field");
643        let mappings = classifier.required_mappings();
644        assert_eq!(mappings.len(), 5);
645        assert!(mappings.iter().any(|m| {
646            m.src_area_id == "field"
647                && m.dst_area_id == "kmem"
648                && m.morphology_id == CLASSIFIER_SCAN_MORPHOLOGY
649        }));
650        assert!(mappings.iter().any(|m| m.src_area_id == "field-b"));
651    }
652
653    #[test]
654    fn deleting_one_field_keeps_the_other_eye() {
655        let mut classifier = sample();
656        classifier
657            .attach_field("field-b".to_string(), "twin-b".to_string())
658            .expect("second field");
659        assert_eq!(classifier.detach_field("field"), Some("twin".to_string()));
660        assert!(classifier.binding_for_field("field").is_none());
661        assert_eq!(
662            classifier
663                .binding_for_field("field-b")
664                .map(|f| f.scan_twin_id.as_str()),
665            Some("twin-b")
666        );
667        assert!(classifier.owns_assembly_core("kmem"));
668        assert!(!classifier.owns_area("twin"));
669        assert!(classifier.owns_area("twin-b"));
670    }
671
672    #[test]
673    fn duplicate_field_mapping_is_rejected() {
674        let mut classifier = sample();
675        let result = classifier.attach_field("field".to_string(), "other-twin".to_string());
676        assert!(result.is_err());
677        assert_eq!(classifier.fields.len(), 1);
678    }
679
680    #[test]
681    fn metadata_update_renames_without_touching_areas() {
682        let mut classifier = sample();
683        classifier
684            .apply_metadata_update(Some("  renamed  ".to_string()), Some([9, 8, 7]))
685            .expect("valid metadata");
686        assert_eq!(classifier.name, "renamed");
687        assert_eq!(classifier.coordinates_3d, [9, 8, 7]);
688        assert_eq!(classifier.kernel_memory_id, "kmem");
689        assert_eq!(classifier.fields[0].scan_twin_id, "twin");
690        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
691    }
692
693    #[test]
694    fn assembly_update_retargets_kernel_and_class_only() {
695        let mut classifier = sample();
696        classifier
697            .apply_assembly_update(
698                None,
699                None,
700                Some("other-region".to_string()),
701                Some("kernel2".to_string()),
702                Some("class2".to_string()),
703            )
704            .expect("valid assembly update");
705        assert_eq!(classifier.parent_region_id, "other-region");
706        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel2"));
707        assert_eq!(classifier.class_area_id.as_deref(), Some("class2"));
708        assert_eq!(classifier.fields[0].field_area_id, "field");
709    }
710
711    #[test]
712    fn scanner_mode_clears_kernel_inputs_and_reports_geometry_change() {
713        let mut classifier = sample();
714        let changed = classifier
715            .apply_training_inputs(
716                ClassifierTrainingMode::Scanner,
717                None,
718                None,
719                Some("mask".to_string()),
720                Some([8, 8, 3]),
721                Some(19),
722            )
723            .expect("scanner inputs");
724        assert!(changed);
725        assert_eq!(classifier.training_mode, ClassifierTrainingMode::Scanner);
726        assert!(classifier.kernel_area_id.is_none());
727        assert!(classifier.class_area_id.is_none());
728        assert_eq!(classifier.mask_area_id.as_deref(), Some("mask"));
729        assert_eq!(classifier.kernel_size, Some([8, 8, 3]));
730        assert_eq!(classifier.class_count, Some(19));
731        assert!(classifier.references_input("mask"));
732        let same = classifier
733            .apply_training_inputs(
734                ClassifierTrainingMode::Scanner,
735                None,
736                None,
737                Some("mask".to_string()),
738                Some([8, 8, 3]),
739                Some(19),
740            )
741            .expect("same scanner geometry");
742        assert!(!same);
743        let recounted = classifier
744            .apply_training_inputs(
745                ClassifierTrainingMode::Scanner,
746                None,
747                None,
748                Some("mask".to_string()),
749                Some([8, 8, 3]),
750                Some(6),
751            )
752            .expect("new class count");
753        assert!(recounted, "a new class count changes what learned ids mean");
754    }
755
756    #[test]
757    fn scanner_mode_requires_a_valid_class_count() {
758        let mut classifier = sample();
759        for bad in [None, Some(0)] {
760            assert!(classifier
761                .apply_training_inputs(
762                    ClassifierTrainingMode::Scanner,
763                    None,
764                    None,
765                    Some("mask".to_string()),
766                    Some([8, 8, 3]),
767                    bad,
768                )
769                .is_err());
770        }
771        assert_eq!(classifier.training_mode, ClassifierTrainingMode::Kernel);
772    }
773
774    #[test]
775    fn kernel_mode_clears_scanner_inputs() {
776        let mut classifier = sample();
777        classifier
778            .apply_training_inputs(
779                ClassifierTrainingMode::Scanner,
780                None,
781                None,
782                Some("mask".to_string()),
783                Some([2, 2, 1]),
784                Some(4),
785            )
786            .expect("scanner");
787        classifier
788            .apply_training_inputs(
789                ClassifierTrainingMode::Kernel,
790                Some("kernel".to_string()),
791                Some("class".to_string()),
792                None,
793                None,
794                None,
795            )
796            .expect("kernel");
797        assert!(classifier.mask_area_id.is_none());
798        assert!(classifier.kernel_size.is_none());
799        assert!(classifier.class_count.is_none());
800        assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
801    }
802
803    #[test]
804    fn detection_twin_is_one_layer_over_the_field() {
805        assert_eq!(detection_twin_shape([256, 128, 3]), [256, 128, 1]);
806    }
807
808    #[test]
809    fn class_output_shape_follows_training_mode() {
810        assert_eq!(
811            class_output_shape(ClassifierTrainingMode::Kernel, [13, 13, 3], 10).unwrap(),
812            [1, 1, 10]
813        );
814        assert_eq!(
815            class_output_shape(ClassifierTrainingMode::Scanner, [13, 13, 3], 10).unwrap(),
816            [13, 13, 1]
817        );
818        assert!(validate_kernel_field([13, 13, 3], [13, 13, 3]).is_ok());
819        assert!(validate_kernel_field([13, 13, 3], [28, 28, 3]).is_err());
820        assert_eq!(
821            class_output_area_name("MNIST classifier", "MNIST kernel", 1),
822            "MNIST classifier class output"
823        );
824        assert_eq!(
825            class_output_area_name("MNIST classifier", "MNIST kernel", 2),
826            "MNIST classifier MNIST kernel class output"
827        );
828        assert!(!class_output_forwards_potential(
829            ClassifierTrainingMode::Kernel
830        ));
831        assert!(class_output_forwards_potential(
832            ClassifierTrainingMode::Scanner
833        ));
834    }
835
836    #[test]
837    fn answer_feedback_matches_class_area_or_detection_output() {
838        assert!(validate_answer_feedback_shape(
839            ClassifierTrainingMode::Kernel,
840            [1, 1, 4],
841            [1, 1, 4],
842            &[],
843            4,
844        )
845        .is_ok());
846        assert!(validate_answer_feedback_shape(
847            ClassifierTrainingMode::Kernel,
848            [8, 8, 4],
849            [1, 1, 4],
850            &[],
851            4,
852        )
853        .is_err());
854        assert!(validate_answer_feedback_shape(
855            ClassifierTrainingMode::Scanner,
856            [16, 16, 1],
857            [16, 16, 1],
858            &[[16, 16, 1]],
859            4,
860        )
861        .is_ok());
862        assert!(validate_answer_feedback_shape(
863            ClassifierTrainingMode::Scanner,
864            [16, 16, 1],
865            [16, 16, 1],
866            &[[16, 16, 1], [8, 8, 1]],
867            4,
868        )
869        .is_err());
870        assert!(
871            validate_answer_feedback_shape(
872                ClassifierTrainingMode::Scanner,
873                [16, 16, 4],
874                [16, 16, 1],
875                &[[16, 16, 1]],
876                4,
877            )
878            .is_err(),
879            "a one-hot class volume no longer matches a potential-coded twin"
880        );
881        assert!(validate_answer_feedback_shape(
882            ClassifierTrainingMode::Scanner,
883            [10, 1, 1],
884            [16, 16, 1],
885            &[[16, 16, 1]],
886            10,
887        )
888        .is_ok());
889    }
890
891    #[test]
892    fn kernel_class_area_must_be_one_by_one_by_n() {
893        assert!(validate_kernel_class_area_shape([1, 1, 4]).is_ok());
894        assert!(validate_kernel_class_area_shape([2, 7, 4]).is_err());
895        assert!(validate_kernel_class_area_shape([4, 1, 1]).is_err());
896        assert!(validate_kernel_class_area_shape([1, 1, 0]).is_err());
897    }
898
899    #[test]
900    fn scanner_field_must_match_mask_and_kernel_depth() {
901        assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 1]).is_ok());
902        assert!(validate_scanner_field([8, 8, 1], [256, 128, 3], [256, 128, 1]).is_err());
903        assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [200, 128, 1]).is_err());
904        assert!(
905            validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 10]).is_err(),
906            "one-hot class masks are replaced by the single-layer potential mask"
907        );
908    }
909}