1use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
21use scirs2_core::numeric::Float;
22use std::fmt::Debug;
23
24use super::{LifelongOptimizer, LifelongStrategy};
25use crate::error::{OptimError, Result};
26use crate::utils::{scalar_or, try_f64};
27
28pub const DEFAULT_TASK_EMBEDDING_DIM: usize = 64;
33
34pub const MAX_TASK_EMBEDDING_DIM: usize = 4096;
37
38pub const DEFAULT_TRANSFER_THRESHOLD: f64 = 0.5;
42
43const WHITENING_EPSILON: f64 = 1e-12;
46
47#[derive(Debug, Clone, Default)]
53pub struct TaskStatistics {
54 observations: usize,
56 mean_gradient: Vec<f64>,
58 mean_squared_gradient: Vec<f64>,
60}
61
62impl TaskStatistics {
63 pub fn observe(&mut self, gradient: &[f64]) {
70 if self.mean_gradient.len() != gradient.len() {
71 self.mean_gradient = vec![0.0; gradient.len()];
72 self.mean_squared_gradient = vec![0.0; gradient.len()];
73 self.observations = 0;
74 }
75
76 self.observations += 1;
77 let weight = 1.0 / self.observations as f64;
78 for ((mean, squared), &value) in self
79 .mean_gradient
80 .iter_mut()
81 .zip(self.mean_squared_gradient.iter_mut())
82 .zip(gradient.iter())
83 {
84 *mean += (value - *mean) * weight;
85 *squared += (value * value - *squared) * weight;
86 }
87 }
88
89 pub fn observations(&self) -> usize {
91 self.observations
92 }
93
94 pub fn mean_gradient(&self) -> &[f64] {
96 &self.mean_gradient
97 }
98
99 pub fn mean_squared_gradient(&self) -> &[f64] {
101 &self.mean_squared_gradient
102 }
103
104 pub fn embedding(&self, dim: usize) -> Vec<f64> {
122 let dim = dim.clamp(1, MAX_TASK_EMBEDDING_DIM);
123 let mut projected = vec![0.0f64; dim];
124
125 for (index, (&mean, &squared)) in self
126 .mean_gradient
127 .iter()
128 .zip(self.mean_squared_gradient.iter())
129 .enumerate()
130 {
131 let whitened = mean / (squared + WHITENING_EPSILON).sqrt();
132 if !whitened.is_finite() {
133 continue;
134 }
135 let hash = mix64(index as u64);
136 let bucket = (hash % dim as u64) as usize;
137 let sign = if (hash >> 63) & 1 == 1 { -1.0 } else { 1.0 };
138 projected[bucket] += sign * whitened;
139 }
140
141 let norm = projected.iter().map(|v| v * v).sum::<f64>().sqrt();
142 if !norm.is_finite() || norm <= 0.0 {
143 return vec![0.0; dim];
144 }
145 for value in projected.iter_mut() {
146 *value /= norm;
147 }
148 projected
149 }
150}
151
152fn mix64(mut x: u64) -> u64 {
160 x ^= x >> 30;
161 x = x.wrapping_mul(0xbf58_476d_1ce4_e5b9);
162 x ^= x >> 27;
163 x = x.wrapping_mul(0x94d0_49bb_1331_11eb);
164 x ^ (x >> 31)
165}
166
167pub fn cosine_similarity(left: &[f64], right: &[f64]) -> f64 {
175 if left.len() != right.len() || left.is_empty() {
176 return 0.0;
177 }
178
179 let mut dot = 0.0;
180 let mut left_norm = 0.0;
181 let mut right_norm = 0.0;
182 for (&a, &b) in left.iter().zip(right.iter()) {
183 dot += a * b;
184 left_norm += a * a;
185 right_norm += b * b;
186 }
187
188 let denominator = (left_norm * right_norm).sqrt();
189 if !denominator.is_finite() || denominator <= 0.0 {
190 return 0.0;
191 }
192 (dot / denominator).clamp(-1.0, 1.0)
193}
194
195fn similarity_key(left: &str, right: &str) -> (String, String) {
198 if left <= right {
199 (left.to_string(), right.to_string())
200 } else {
201 (right.to_string(), left.to_string())
202 }
203}
204
205#[derive(Debug, Clone, PartialEq)]
207pub struct TransferOutcome {
208 pub source_task: Option<String>,
210 pub similarity: f64,
213 pub transfer_weight: f64,
216}
217
218impl<A: Float + ScalarOperand + Debug + std::iter::Sum, D: Dimension + Send + Sync>
219 LifelongOptimizer<A, D>
220{
221 pub fn transfer_threshold(&self) -> f64 {
224 self.transfer_threshold
225 }
226
227 pub fn set_transfer_threshold(&mut self, threshold: f64) -> Result<()> {
233 if !threshold.is_finite() || !(-1.0..=1.0).contains(&threshold) {
234 return Err(OptimError::InvalidConfig(format!(
235 "transfer threshold {threshold} is not a cosine similarity in [-1, 1]"
236 )));
237 }
238 self.transfer_threshold = threshold;
239 self.rebuild_task_clusters();
240 Ok(())
241 }
242
243 pub fn task_embedding(&self, task_id: &str) -> Option<&[f64]> {
245 self.shared_knowledge
246 .task_embeddings
247 .get(task_id)
248 .map(|embedding| embedding.as_slice())
249 }
250
251 pub fn task_statistics(&self, task_id: &str) -> Option<&TaskStatistics> {
253 self.shared_knowledge.task_statistics.get(task_id)
254 }
255
256 pub fn transfer_weight(&self, source: &str, target: &str) -> f64 {
259 self.shared_knowledge
260 .transfer_weights
261 .get(&(source.to_string(), target.to_string()))
262 .copied()
263 .unwrap_or(0.0)
264 }
265
266 pub fn task_dependencies(&self, task_id: &str) -> &[String] {
268 self.task_graph
269 .task_dependencies
270 .get(task_id)
271 .map(|dependencies| dependencies.as_slice())
272 .unwrap_or(&[])
273 }
274
275 pub fn task_clusters(&self) -> &[Vec<String>] {
278 &self.task_graph.task_clusters
279 }
280
281 pub fn task_parameters(&self, task_id: &str) -> Option<&Array<A, D>> {
283 self.task_optimizers
284 .get(task_id)
285 .map(|optimizer| optimizer.parameters())
286 }
287
288 pub fn mean_transfer_weight(&self) -> f64 {
293 let weights = &self.shared_knowledge.transfer_weights;
294 if weights.is_empty() {
295 return 0.0;
296 }
297 weights.values().sum::<f64>() / weights.len() as f64
298 }
299
300 pub fn start_task_with_probe(
322 &mut self,
323 task_id: String,
324 initial_parameters: Array<A, D>,
325 probe_gradient: &Array<A, D>,
326 ) -> Result<TransferOutcome> {
327 if probe_gradient.raw_dim() != initial_parameters.raw_dim() {
328 return Err(OptimError::DimensionMismatch(format!(
329 "transfer probe: initial parameters have shape {:?} but the probe \
330 gradient has shape {:?}",
331 initial_parameters.raw_dim().slice(),
332 probe_gradient.raw_dim().slice()
333 )));
334 }
335
336 let dim = self.embedding_dim();
337 let probe_values = flatten_to_f64(probe_gradient)?;
338 let mut probe_statistics = TaskStatistics::default();
339 probe_statistics.observe(&probe_values);
340 let probe_embedding = probe_statistics.embedding(dim);
341
342 let mut best: Option<(String, f64)> = None;
346 for (candidate, embedding) in &self.shared_knowledge.task_embeddings {
347 if *candidate == task_id {
348 continue;
349 }
350 let similarity = cosine_similarity(&probe_embedding, embedding);
351 self.task_graph
352 .task_similarities
353 .insert(similarity_key(&task_id, candidate), similarity);
354
355 let better = match &best {
358 None => true,
359 Some((best_id, best_similarity)) => {
360 similarity > *best_similarity
361 || (similarity == *best_similarity && candidate < best_id)
362 }
363 };
364 if better {
365 best = Some((candidate.clone(), similarity));
366 }
367 }
368
369 let mut parameters = initial_parameters;
370 let mut outcome = TransferOutcome {
371 source_task: None,
372 similarity: best.as_ref().map(|(_, s)| *s).unwrap_or(0.0),
373 transfer_weight: 0.0,
374 };
375
376 if let Some((source, similarity)) = best {
377 if similarity >= self.transfer_threshold {
378 let source_parameters = self.task_parameters(&source).cloned();
379 if let Some(source_parameters) = source_parameters {
380 if source_parameters.raw_dim() == parameters.raw_dim() {
381 let weight = similarity.clamp(0.0, 1.0);
382 let blend = scalar_or(weight, A::zero());
383 let keep = A::one() - blend;
384 for (slot, &learned) in parameters.iter_mut().zip(source_parameters.iter())
385 {
386 *slot = *slot * keep + learned * blend;
387 }
388 outcome.source_task = Some(source.clone());
389 outcome.transfer_weight = weight;
390 }
391 }
392 }
393 }
394
395 self.start_task(task_id.clone(), parameters)?;
396
397 if let Some(source) = outcome.source_task.clone() {
398 self.shared_knowledge
399 .transfer_weights
400 .insert((source.clone(), task_id.clone()), outcome.transfer_weight);
401 let dependencies = self
402 .task_graph
403 .task_dependencies
404 .entry(task_id.clone())
405 .or_default();
406 if !dependencies.contains(&source) {
407 dependencies.push(source);
408 }
409 }
410
411 self.shared_knowledge
414 .task_statistics
415 .insert(task_id.clone(), probe_statistics);
416 self.shared_knowledge
417 .task_embeddings
418 .insert(task_id, probe_embedding);
419 self.rebuild_task_clusters();
420
421 Ok(outcome)
422 }
423
424 pub(super) fn record_task_observation(&mut self, gradient: &Array<A, D>) -> Result<()> {
431 let Some(task_id) = self.current_task.clone() else {
432 return Ok(());
433 };
434
435 let dim = self.embedding_dim();
436 let values = flatten_to_f64(gradient)?;
437
438 let statistics = self
439 .shared_knowledge
440 .task_statistics
441 .entry(task_id.clone())
442 .or_default();
443 statistics.observe(&values);
444 let embedding = statistics.embedding(dim);
445
446 self.shared_knowledge
447 .task_embeddings
448 .insert(task_id.clone(), embedding);
449 self.refresh_similarities(&task_id);
450 self.rebuild_task_clusters();
451
452 Ok(())
453 }
454
455 fn embedding_dim(&self) -> usize {
458 match self.strategy {
459 LifelongStrategy::MetaLearning {
460 task_embedding_size,
461 ..
462 } => task_embedding_size.clamp(1, MAX_TASK_EMBEDDING_DIM),
463 _ => DEFAULT_TASK_EMBEDDING_DIM,
464 }
465 }
466
467 fn refresh_similarities(&mut self, task_id: &str) {
469 let Some(embedding) = self.shared_knowledge.task_embeddings.get(task_id).cloned() else {
470 return;
471 };
472
473 for (other, other_embedding) in &self.shared_knowledge.task_embeddings {
474 if other == task_id {
475 continue;
476 }
477 let similarity = cosine_similarity(&embedding, other_embedding);
478 self.task_graph
479 .task_similarities
480 .insert(similarity_key(task_id, other), similarity);
481 }
482 }
483
484 fn rebuild_task_clusters(&mut self) {
501 let mut tasks: Vec<String> = self
502 .shared_knowledge
503 .task_embeddings
504 .keys()
505 .cloned()
506 .collect();
507 tasks.sort();
508
509 let count = tasks.len();
510 if count == 0 {
511 self.task_graph.task_clusters.clear();
512 return;
513 }
514
515 let threshold = self.transfer_threshold;
516 let mut parent: Vec<usize> = (0..count).collect();
517 for left in 0..count {
518 for right in (left + 1)..count {
519 if self.compute_task_similarity(&tasks[left], &tasks[right]) >= threshold {
520 let left_root = find_root(&mut parent, left);
521 let right_root = find_root(&mut parent, right);
522 if left_root != right_root {
523 parent[left_root.max(right_root)] = left_root.min(right_root);
524 }
525 }
526 }
527 }
528
529 let mut buckets: Vec<Vec<String>> = vec![Vec::new(); count];
530 for (index, task) in tasks.iter().enumerate() {
531 let root = find_root(&mut parent, index);
532 buckets[root].push(task.clone());
533 }
534
535 let mut named: Vec<Vec<String>> = buckets
536 .into_iter()
537 .filter(|cluster| !cluster.is_empty())
538 .collect();
539 named.sort();
540 self.task_graph.task_clusters = named;
541 }
542}
543
544fn find_root(parent: &mut [usize], node: usize) -> usize {
546 let mut current = node;
547 while parent[current] != current {
548 parent[current] = parent[parent[current]];
549 current = parent[current];
550 }
551 current
552}
553
554fn flatten_to_f64<A: Float, D: Dimension>(array: &Array<A, D>) -> Result<Vec<f64>> {
557 array.iter().map(|&value| try_f64(value)).collect()
558}
559
560#[cfg(test)]
561mod tests {
562 use super::*;
563 use crate::online_learning::MemoryUpdateStrategy;
564 use scirs2_core::ndarray::{Array1, Ix1};
565
566 fn quadratic_gradient(parameters: &Array1<f64>, target: &Array1<f64>) -> Array1<f64> {
569 parameters
570 .iter()
571 .zip(target.iter())
572 .map(|(&x, &t)| 2.0 * (x - t))
573 .collect()
574 }
575
576 fn quadratic_loss(parameters: &Array1<f64>, target: &Array1<f64>) -> f64 {
577 parameters
578 .iter()
579 .zip(target.iter())
580 .map(|(&x, &t)| (x - t) * (x - t))
581 .sum()
582 }
583
584 fn optimizer() -> LifelongOptimizer<f64, Ix1> {
585 LifelongOptimizer::new(LifelongStrategy::MemoryAugmented {
586 memory_size: 128,
587 update_strategy: MemoryUpdateStrategy::FIFO,
588 })
589 }
590
591 fn train(
593 optimizer: &mut LifelongOptimizer<f64, Ix1>,
594 task_id: &str,
595 target: &Array1<f64>,
596 steps: usize,
597 ) {
598 for _ in 0..steps {
599 let parameters = optimizer
600 .task_parameters(task_id)
601 .expect("task must exist")
602 .clone();
603 let gradient = quadratic_gradient(¶meters, target);
604 let loss = quadratic_loss(¶meters, target);
605 optimizer
606 .update_current_task(&gradient, loss)
607 .expect("update must succeed");
608 }
609 }
610
611 #[test]
614 fn a_similar_task_is_warm_started_and_drops_its_loss_faster() {
615 let target_a = Array1::from_vec(vec![3.0, 3.0, 3.0, 3.0]);
616 let target_b = Array1::from_vec(vec![3.05, 3.05, 3.05, 3.05]);
617 let cold_start = Array1::zeros(4);
618
619 let mut opt = optimizer();
620 opt.start_task("a".to_string(), cold_start.clone())
621 .expect("start a");
622 train(&mut opt, "a", &target_a, 4000);
623
624 let learned_a = opt.task_parameters("a").expect("task a").clone();
625 assert!(
626 quadratic_loss(&learned_a, &target_a) < 1.0,
627 "task a did not learn: {learned_a:?}"
628 );
629
630 let probe = quadratic_gradient(&cold_start, &target_b);
632 let outcome = opt
633 .start_task_with_probe("b".to_string(), cold_start.clone(), &probe)
634 .expect("start b");
635
636 assert_eq!(
637 outcome.source_task.as_deref(),
638 Some("a"),
639 "a task pulling the same way must transfer (similarity {})",
640 outcome.similarity
641 );
642 assert!(
643 outcome.similarity > 0.9,
644 "similarity {} is implausibly low for near-identical tasks",
645 outcome.similarity
646 );
647 assert!(outcome.transfer_weight > 0.9);
648
649 let warm_parameters = opt.task_parameters("b").expect("task b").clone();
650 let warm_initial_loss = quadratic_loss(&warm_parameters, &target_b);
651 let cold_initial_loss = quadratic_loss(&cold_start, &target_b);
652 assert!(
653 warm_initial_loss < cold_initial_loss * 0.1,
654 "warm start did not help: {warm_initial_loss} vs cold {cold_initial_loss}"
655 );
656
657 train(&mut opt, "b", &target_b, 20);
660 let warm_after = quadratic_loss(opt.task_parameters("b").expect("task b"), &target_b);
661
662 let mut cold = optimizer();
663 cold.start_task("b".to_string(), cold_start.clone())
664 .expect("cold start b");
665 train(&mut cold, "b", &target_b, 20);
666 let cold_after = quadratic_loss(cold.task_parameters("b").expect("cold task b"), &target_b);
667
668 assert!(
669 warm_after < cold_after,
670 "warm start ({warm_after}) must beat cold start ({cold_after}) after 20 steps"
671 );
672 }
673
674 #[test]
677 fn a_dissimilar_task_is_not_warm_started() {
678 let target_a = Array1::from_vec(vec![3.0, 3.0, 3.0, 3.0]);
679 let target_c = Array1::from_vec(vec![-3.0, -3.0, -3.0, -3.0]);
680 let cold_start = Array1::zeros(4);
681
682 let mut opt = optimizer();
683 opt.start_task("a".to_string(), cold_start.clone())
684 .expect("start a");
685 train(&mut opt, "a", &target_a, 500);
686
687 let probe = quadratic_gradient(&cold_start, &target_c);
688 let outcome = opt
689 .start_task_with_probe("c".to_string(), cold_start.clone(), &probe)
690 .expect("start c");
691
692 assert_eq!(
693 outcome.source_task, None,
694 "an opposing task must not transfer (similarity {})",
695 outcome.similarity
696 );
697 assert!(
698 outcome.similarity < 0.0,
699 "opposing gradients must score below zero, got {}",
700 outcome.similarity
701 );
702 assert_eq!(outcome.transfer_weight, 0.0);
703 assert_eq!(
704 opt.task_parameters("c").expect("task c"),
705 &cold_start,
706 "a task that did not transfer must start exactly where the caller put it"
707 );
708 assert!(opt.task_dependencies("c").is_empty());
709 }
710
711 #[test]
714 fn transfer_weights_and_dependencies_are_recorded() {
715 let target_a = Array1::from_vec(vec![2.0, -2.0]);
716 let target_b = Array1::from_vec(vec![2.1, -2.1]);
717 let cold_start = Array1::zeros(2);
718
719 let mut opt = optimizer();
720 opt.start_task("a".to_string(), cold_start.clone())
721 .expect("start a");
722 train(&mut opt, "a", &target_a, 500);
723
724 let probe = quadratic_gradient(&cold_start, &target_b);
725 opt.start_task_with_probe("b".to_string(), cold_start, &probe)
726 .expect("start b");
727
728 assert_eq!(opt.task_dependencies("b").to_vec(), vec!["a".to_string()]);
729 assert!(opt.transfer_weight("a", "b") > 0.9);
730 assert_eq!(
731 opt.transfer_weight("b", "a"),
732 0.0,
733 "transfer is directed: b was started from a, not the other way round"
734 );
735 assert!(
736 opt.compute_task_similarity("a", "b") > 0.9,
737 "the similarity matrix must be populated"
738 );
739 assert!(
740 opt.get_lifelong_stats().transfer_efficiency > 0.9,
741 "transfer efficiency must reflect the transfers that happened"
742 );
743 assert!(opt.task_embedding("a").is_some());
744 assert!(opt.task_embedding("b").is_some());
745 }
746
747 #[test]
749 fn similar_tasks_cluster_together() {
750 let cold_start = Array1::zeros(3);
751 let target_a = Array1::from_vec(vec![1.0, 1.0, 1.0]);
752 let target_b = Array1::from_vec(vec![1.2, 1.1, 1.05]);
753 let target_c = Array1::from_vec(vec![-1.0, -1.0, -1.0]);
754
755 let mut opt = optimizer();
756 opt.start_task("a".to_string(), cold_start.clone())
757 .expect("start a");
758 train(&mut opt, "a", &target_a, 200);
759
760 opt.start_task("b".to_string(), cold_start.clone())
761 .expect("start b");
762 train(&mut opt, "b", &target_b, 200);
763
764 opt.start_task("c".to_string(), cold_start)
765 .expect("start c");
766 train(&mut opt, "c", &target_c, 200);
767
768 let clusters = opt.task_clusters().to_vec();
769 assert_eq!(
770 clusters,
771 vec![
772 vec!["a".to_string(), "b".to_string()],
773 vec!["c".to_string()]
774 ],
775 "clusters: {clusters:?}, sim(a,b) = {}, sim(a,c) = {}",
776 opt.compute_task_similarity("a", "b"),
777 opt.compute_task_similarity("a", "c")
778 );
779 }
780
781 #[test]
792 fn the_threshold_controls_clustering_and_transfer() {
793 let cold_start = Array1::zeros(6);
794 let target_a = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0]);
795 let target_b = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0, -1.0]);
796
797 let mut opt = optimizer();
798 opt.start_task("a".to_string(), cold_start.clone())
799 .expect("start a");
800 train(&mut opt, "a", &target_a, 200);
801 opt.start_task("b".to_string(), cold_start.clone())
802 .expect("start b");
803 train(&mut opt, "b", &target_b, 200);
804
805 let similarity = opt.compute_task_similarity("a", "b");
806 assert!(
807 (0.5..0.9).contains(&similarity),
808 "five agreeing coordinates out of six should score in [0.5, 0.9), got {similarity}"
809 );
810 assert_eq!(
811 opt.task_clusters().len(),
812 1,
813 "similar tasks share a cluster"
814 );
815
816 opt.set_transfer_threshold(0.9).expect("valid threshold");
817 assert_eq!(
818 opt.task_clusters().len(),
819 2,
820 "a threshold above the measured similarity must separate the tasks"
821 );
822
823 let target_d = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, -1.0, 1.0]);
827 let probe = quadratic_gradient(&cold_start, &target_d);
828 let outcome = opt
829 .start_task_with_probe("d".to_string(), cold_start, &probe)
830 .expect("start d");
831 assert!(
832 (0.5..0.9).contains(&outcome.similarity),
833 "probe similarity {} is outside the band this test needs",
834 outcome.similarity
835 );
836 assert_eq!(
837 outcome.source_task, None,
838 "similarity {} cleared a threshold of 0.9",
839 outcome.similarity
840 );
841
842 assert!(opt.set_transfer_threshold(2.0).is_err());
843 assert!(opt.set_transfer_threshold(f64::NAN).is_err());
844 }
845
846 #[test]
848 fn a_mismatched_probe_is_reported() {
849 let mut opt = optimizer();
850 let error = opt
851 .start_task_with_probe(
852 "a".to_string(),
853 Array1::zeros(3),
854 &Array1::from_vec(vec![1.0, 2.0]),
855 )
856 .expect_err("a shape mismatch must be reported");
857 assert!(
858 matches!(error, OptimError::DimensionMismatch(_)),
859 "{error:?}"
860 );
861 }
862
863 #[test]
865 fn the_first_task_has_nothing_to_transfer_from() {
866 let mut opt = optimizer();
867 let outcome = opt
868 .start_task_with_probe(
869 "a".to_string(),
870 Array1::zeros(2),
871 &Array1::from_vec(vec![1.0, -1.0]),
872 )
873 .expect("start a");
874 assert_eq!(outcome.source_task, None);
875 assert_eq!(outcome.similarity, 0.0);
876 assert_eq!(opt.mean_transfer_weight(), 0.0);
877 assert_eq!(opt.task_clusters().to_vec(), vec![vec!["a".to_string()]]);
878 }
879
880 #[test]
883 fn embeddings_have_a_fixed_width_regardless_of_parameter_count() {
884 let mut small = TaskStatistics::default();
885 small.observe(&[1.0, -1.0, 1.0]);
886 let mut large = TaskStatistics::default();
887 large.observe(&vec![0.5; 500]);
888
889 assert_eq!(small.embedding(64).len(), 64);
890 assert_eq!(large.embedding(64).len(), 64);
891 let similarity = cosine_similarity(&small.embedding(64), &large.embedding(64));
892 assert!(
893 (-1.0..=1.0).contains(&similarity),
894 "similarity {similarity} is not a cosine"
895 );
896 }
897
898 #[test]
901 fn the_embedding_captures_gradient_direction() {
902 let mut forward = TaskStatistics::default();
903 let mut backward = TaskStatistics::default();
904 for step in 0..10 {
905 let scale = 1.0 + step as f64;
906 forward.observe(&[scale, 2.0 * scale, -scale]);
907 backward.observe(&[-scale, -2.0 * scale, scale]);
908 }
909
910 let same = cosine_similarity(&forward.embedding(32), &forward.embedding(32));
911 let opposite = cosine_similarity(&forward.embedding(32), &backward.embedding(32));
912 assert!((same - 1.0).abs() < 1e-9, "same direction scored {same}");
913 assert!(
914 (opposite + 1.0).abs() < 1e-9,
915 "opposite direction scored {opposite}"
916 );
917 }
918
919 #[test]
922 fn a_directionless_task_has_a_zero_embedding() {
923 let mut statistics = TaskStatistics::default();
924 statistics.observe(&[1.0, 1.0]);
925 statistics.observe(&[-1.0, -1.0]);
926 let embedding = statistics.embedding(8);
927 assert!(embedding.iter().all(|&value| value == 0.0));
928 assert_eq!(cosine_similarity(&embedding, &embedding), 0.0);
929 }
930
931 #[test]
934 fn forgetting_is_measured_from_re_evaluated_tasks() {
935 let cold_start = Array1::zeros(2);
936 let target_a = Array1::from_vec(vec![1.0, 1.0]);
937
938 let mut opt = optimizer();
939 opt.start_task("a".to_string(), cold_start.clone())
940 .expect("start a");
941 train(&mut opt, "a", &target_a, 50);
942
943 assert_eq!(
944 opt.get_lifelong_stats().catastrophic_forgetting,
945 0.0,
946 "forgetting is not observable while the task is still current"
947 );
948
949 opt.start_task("b".to_string(), cold_start)
950 .expect("start b");
951 assert_eq!(
952 opt.get_lifelong_stats().catastrophic_forgetting,
953 0.0,
954 "switching away does not by itself demonstrate forgetting"
955 );
956
957 let reference = opt
959 .task_performance
960 .get("a")
961 .and_then(|history| history.last())
962 .copied()
963 .expect("task a recorded losses");
964 opt.record_task_performance("a", reference + 0.25)
965 .expect("recording a re-evaluation must succeed");
966 let forgetting = opt.get_lifelong_stats().catastrophic_forgetting;
967 assert!(
968 (forgetting - 0.25).abs() < 1e-9,
969 "forgetting was not measured from the re-evaluation: {forgetting}"
970 );
971
972 assert!(opt.record_task_performance("b", 1.0).is_err());
974 assert!(opt.record_task_performance("nope", 1.0).is_err());
975 }
976
977 #[test]
979 fn statistics_are_running_means() {
980 let mut statistics = TaskStatistics::default();
981 statistics.observe(&[2.0, 0.0]);
982 statistics.observe(&[4.0, 0.0]);
983 assert_eq!(statistics.observations(), 2);
984 assert!((statistics.mean_gradient()[0] - 3.0).abs() < 1e-12);
985 assert!((statistics.mean_squared_gradient()[0] - 10.0).abs() < 1e-12);
986 }
987}