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_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
469pub 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
481pub 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}