1use crate::neuron_voxels::class_potential::validate_class_count;
18use serde::{Deserialize, Serialize};
19use std::collections::HashMap;
20
21pub const CLASSIFIER_KERNEL_MORPHOLOGY: &str = "episodic_memory";
23pub const CLASSIFIER_CLASS_MORPHOLOGY: &str = "episodic_memory";
25pub const CLASSIFIER_ASSOCIATIVE_MORPHOLOGY: &str = "associative_memory";
27pub const CLASSIFIER_SCAN_MORPHOLOGY: &str = "episodic_scan";
29
30#[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#[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#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)]
47#[serde(rename_all = "snake_case")]
48pub enum ClassifierTrainingMode {
49 #[default]
51 Kernel,
52 Scanner,
54}
55
56impl ClassifierTrainingMode {
57 pub fn as_str(self) -> &'static str {
59 match self {
60 Self::Kernel => "kernel",
61 Self::Scanner => "scanner",
62 }
63 }
64
65 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
79pub struct Classifier {
80 pub classifier_id: String,
82 pub name: String,
83 pub parent_region_id: String,
85 pub coordinates_3d: [i32; 3],
86 #[serde(default)]
89 pub training_mode: ClassifierTrainingMode,
90 #[serde(default, skip_serializing_if = "Option::is_none")]
92 pub kernel_area_id: Option<String>,
93 #[serde(default, skip_serializing_if = "Option::is_none")]
95 pub class_area_id: Option<String>,
96 #[serde(default, skip_serializing_if = "Option::is_none")]
99 pub mask_area_id: Option<String>,
100 #[serde(default, skip_serializing_if = "Option::is_none")]
102 pub class_count: Option<u32>,
103 #[serde(default, skip_serializing_if = "Option::is_none")]
105 pub kernel_size: Option<[u32; 3]>,
106 #[serde(default)]
108 pub fields: Vec<ClassifierField>,
109 pub kernel_memory_id: String,
111 pub class_memory_id: String,
112 #[serde(default)]
115 pub reward_training: bool,
116 #[serde(default, skip_serializing_if = "Option::is_none")]
119 pub answer_feedback_area_id: Option<String>,
120 #[serde(default, skip_serializing_if = "Option::is_none")]
122 pub pain_area_id: Option<String>,
123 #[serde(default, skip_serializing_if = "Option::is_none")]
125 pub pleasure_area_id: Option<String>,
126 #[serde(default)]
128 pub answer_latency_bursts: u32,
129 #[serde(default, skip_serializing_if = "Option::is_none")]
131 pub learn_area_id: Option<String>,
132 #[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 pub fn assembly_core_ids(&self) -> Vec<String> {
142 vec![self.kernel_memory_id.clone(), self.class_memory_id.clone()]
143 }
144
145 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 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 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 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 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 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 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 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 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 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 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 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
421pub 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
431pub 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
439pub fn detection_twin_shape(field: [u32; 3]) -> [u32; 3] {
441 [field[0], field[1], 1]
442}
443
444pub 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
458pub 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
481pub fn class_output_forwards_potential(mode: ClassifierTrainingMode) -> bool {
483 matches!(mode, ClassifierTrainingMode::Scanner)
484}
485
486pub 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
499pub 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
524pub 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
536pub 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
577pub 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}