use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
use scirs2_core::numeric::Float;
use std::fmt::Debug;
use super::{LifelongOptimizer, LifelongStrategy};
use crate::error::{OptimError, Result};
use crate::utils::{scalar_or, try_f64};
pub const DEFAULT_TASK_EMBEDDING_DIM: usize = 64;
pub const MAX_TASK_EMBEDDING_DIM: usize = 4096;
pub const DEFAULT_TRANSFER_THRESHOLD: f64 = 0.5;
const WHITENING_EPSILON: f64 = 1e-12;
#[derive(Debug, Clone, Default)]
pub struct TaskStatistics {
observations: usize,
mean_gradient: Vec<f64>,
mean_squared_gradient: Vec<f64>,
}
impl TaskStatistics {
pub fn observe(&mut self, gradient: &[f64]) {
if self.mean_gradient.len() != gradient.len() {
self.mean_gradient = vec![0.0; gradient.len()];
self.mean_squared_gradient = vec![0.0; gradient.len()];
self.observations = 0;
}
self.observations += 1;
let weight = 1.0 / self.observations as f64;
for ((mean, squared), &value) in self
.mean_gradient
.iter_mut()
.zip(self.mean_squared_gradient.iter_mut())
.zip(gradient.iter())
{
*mean += (value - *mean) * weight;
*squared += (value * value - *squared) * weight;
}
}
pub fn observations(&self) -> usize {
self.observations
}
pub fn mean_gradient(&self) -> &[f64] {
&self.mean_gradient
}
pub fn mean_squared_gradient(&self) -> &[f64] {
&self.mean_squared_gradient
}
pub fn embedding(&self, dim: usize) -> Vec<f64> {
let dim = dim.clamp(1, MAX_TASK_EMBEDDING_DIM);
let mut projected = vec![0.0f64; dim];
for (index, (&mean, &squared)) in self
.mean_gradient
.iter()
.zip(self.mean_squared_gradient.iter())
.enumerate()
{
let whitened = mean / (squared + WHITENING_EPSILON).sqrt();
if !whitened.is_finite() {
continue;
}
let hash = mix64(index as u64);
let bucket = (hash % dim as u64) as usize;
let sign = if (hash >> 63) & 1 == 1 { -1.0 } else { 1.0 };
projected[bucket] += sign * whitened;
}
let norm = projected.iter().map(|v| v * v).sum::<f64>().sqrt();
if !norm.is_finite() || norm <= 0.0 {
return vec![0.0; dim];
}
for value in projected.iter_mut() {
*value /= norm;
}
projected
}
}
fn mix64(mut x: u64) -> u64 {
x ^= x >> 30;
x = x.wrapping_mul(0xbf58_476d_1ce4_e5b9);
x ^= x >> 27;
x = x.wrapping_mul(0x94d0_49bb_1331_11eb);
x ^ (x >> 31)
}
pub fn cosine_similarity(left: &[f64], right: &[f64]) -> f64 {
if left.len() != right.len() || left.is_empty() {
return 0.0;
}
let mut dot = 0.0;
let mut left_norm = 0.0;
let mut right_norm = 0.0;
for (&a, &b) in left.iter().zip(right.iter()) {
dot += a * b;
left_norm += a * a;
right_norm += b * b;
}
let denominator = (left_norm * right_norm).sqrt();
if !denominator.is_finite() || denominator <= 0.0 {
return 0.0;
}
(dot / denominator).clamp(-1.0, 1.0)
}
fn similarity_key(left: &str, right: &str) -> (String, String) {
if left <= right {
(left.to_string(), right.to_string())
} else {
(right.to_string(), left.to_string())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TransferOutcome {
pub source_task: Option<String>,
pub similarity: f64,
pub transfer_weight: f64,
}
impl<A: Float + ScalarOperand + Debug + std::iter::Sum, D: Dimension + Send + Sync>
LifelongOptimizer<A, D>
{
pub fn transfer_threshold(&self) -> f64 {
self.transfer_threshold
}
pub fn set_transfer_threshold(&mut self, threshold: f64) -> Result<()> {
if !threshold.is_finite() || !(-1.0..=1.0).contains(&threshold) {
return Err(OptimError::InvalidConfig(format!(
"transfer threshold {threshold} is not a cosine similarity in [-1, 1]"
)));
}
self.transfer_threshold = threshold;
self.rebuild_task_clusters();
Ok(())
}
pub fn task_embedding(&self, task_id: &str) -> Option<&[f64]> {
self.shared_knowledge
.task_embeddings
.get(task_id)
.map(|embedding| embedding.as_slice())
}
pub fn task_statistics(&self, task_id: &str) -> Option<&TaskStatistics> {
self.shared_knowledge.task_statistics.get(task_id)
}
pub fn transfer_weight(&self, source: &str, target: &str) -> f64 {
self.shared_knowledge
.transfer_weights
.get(&(source.to_string(), target.to_string()))
.copied()
.unwrap_or(0.0)
}
pub fn task_dependencies(&self, task_id: &str) -> &[String] {
self.task_graph
.task_dependencies
.get(task_id)
.map(|dependencies| dependencies.as_slice())
.unwrap_or(&[])
}
pub fn task_clusters(&self) -> &[Vec<String>] {
&self.task_graph.task_clusters
}
pub fn task_parameters(&self, task_id: &str) -> Option<&Array<A, D>> {
self.task_optimizers
.get(task_id)
.map(|optimizer| optimizer.parameters())
}
pub fn mean_transfer_weight(&self) -> f64 {
let weights = &self.shared_knowledge.transfer_weights;
if weights.is_empty() {
return 0.0;
}
weights.values().sum::<f64>() / weights.len() as f64
}
pub fn start_task_with_probe(
&mut self,
task_id: String,
initial_parameters: Array<A, D>,
probe_gradient: &Array<A, D>,
) -> Result<TransferOutcome> {
if probe_gradient.raw_dim() != initial_parameters.raw_dim() {
return Err(OptimError::DimensionMismatch(format!(
"transfer probe: initial parameters have shape {:?} but the probe \
gradient has shape {:?}",
initial_parameters.raw_dim().slice(),
probe_gradient.raw_dim().slice()
)));
}
let dim = self.embedding_dim();
let probe_values = flatten_to_f64(probe_gradient)?;
let mut probe_statistics = TaskStatistics::default();
probe_statistics.observe(&probe_values);
let probe_embedding = probe_statistics.embedding(dim);
let mut best: Option<(String, f64)> = None;
for (candidate, embedding) in &self.shared_knowledge.task_embeddings {
if *candidate == task_id {
continue;
}
let similarity = cosine_similarity(&probe_embedding, embedding);
self.task_graph
.task_similarities
.insert(similarity_key(&task_id, candidate), similarity);
let better = match &best {
None => true,
Some((best_id, best_similarity)) => {
similarity > *best_similarity
|| (similarity == *best_similarity && candidate < best_id)
}
};
if better {
best = Some((candidate.clone(), similarity));
}
}
let mut parameters = initial_parameters;
let mut outcome = TransferOutcome {
source_task: None,
similarity: best.as_ref().map(|(_, s)| *s).unwrap_or(0.0),
transfer_weight: 0.0,
};
if let Some((source, similarity)) = best {
if similarity >= self.transfer_threshold {
let source_parameters = self.task_parameters(&source).cloned();
if let Some(source_parameters) = source_parameters {
if source_parameters.raw_dim() == parameters.raw_dim() {
let weight = similarity.clamp(0.0, 1.0);
let blend = scalar_or(weight, A::zero());
let keep = A::one() - blend;
for (slot, &learned) in parameters.iter_mut().zip(source_parameters.iter())
{
*slot = *slot * keep + learned * blend;
}
outcome.source_task = Some(source.clone());
outcome.transfer_weight = weight;
}
}
}
}
self.start_task(task_id.clone(), parameters)?;
if let Some(source) = outcome.source_task.clone() {
self.shared_knowledge
.transfer_weights
.insert((source.clone(), task_id.clone()), outcome.transfer_weight);
let dependencies = self
.task_graph
.task_dependencies
.entry(task_id.clone())
.or_default();
if !dependencies.contains(&source) {
dependencies.push(source);
}
}
self.shared_knowledge
.task_statistics
.insert(task_id.clone(), probe_statistics);
self.shared_knowledge
.task_embeddings
.insert(task_id, probe_embedding);
self.rebuild_task_clusters();
Ok(outcome)
}
pub(super) fn record_task_observation(&mut self, gradient: &Array<A, D>) -> Result<()> {
let Some(task_id) = self.current_task.clone() else {
return Ok(());
};
let dim = self.embedding_dim();
let values = flatten_to_f64(gradient)?;
let statistics = self
.shared_knowledge
.task_statistics
.entry(task_id.clone())
.or_default();
statistics.observe(&values);
let embedding = statistics.embedding(dim);
self.shared_knowledge
.task_embeddings
.insert(task_id.clone(), embedding);
self.refresh_similarities(&task_id);
self.rebuild_task_clusters();
Ok(())
}
fn embedding_dim(&self) -> usize {
match self.strategy {
LifelongStrategy::MetaLearning {
task_embedding_size,
..
} => task_embedding_size.clamp(1, MAX_TASK_EMBEDDING_DIM),
_ => DEFAULT_TASK_EMBEDDING_DIM,
}
}
fn refresh_similarities(&mut self, task_id: &str) {
let Some(embedding) = self.shared_knowledge.task_embeddings.get(task_id).cloned() else {
return;
};
for (other, other_embedding) in &self.shared_knowledge.task_embeddings {
if other == task_id {
continue;
}
let similarity = cosine_similarity(&embedding, other_embedding);
self.task_graph
.task_similarities
.insert(similarity_key(task_id, other), similarity);
}
}
fn rebuild_task_clusters(&mut self) {
let mut tasks: Vec<String> = self
.shared_knowledge
.task_embeddings
.keys()
.cloned()
.collect();
tasks.sort();
let count = tasks.len();
if count == 0 {
self.task_graph.task_clusters.clear();
return;
}
let threshold = self.transfer_threshold;
let mut parent: Vec<usize> = (0..count).collect();
for left in 0..count {
for right in (left + 1)..count {
if self.compute_task_similarity(&tasks[left], &tasks[right]) >= threshold {
let left_root = find_root(&mut parent, left);
let right_root = find_root(&mut parent, right);
if left_root != right_root {
parent[left_root.max(right_root)] = left_root.min(right_root);
}
}
}
}
let mut buckets: Vec<Vec<String>> = vec![Vec::new(); count];
for (index, task) in tasks.iter().enumerate() {
let root = find_root(&mut parent, index);
buckets[root].push(task.clone());
}
let mut named: Vec<Vec<String>> = buckets
.into_iter()
.filter(|cluster| !cluster.is_empty())
.collect();
named.sort();
self.task_graph.task_clusters = named;
}
}
fn find_root(parent: &mut [usize], node: usize) -> usize {
let mut current = node;
while parent[current] != current {
parent[current] = parent[parent[current]];
current = parent[current];
}
current
}
fn flatten_to_f64<A: Float, D: Dimension>(array: &Array<A, D>) -> Result<Vec<f64>> {
array.iter().map(|&value| try_f64(value)).collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::online_learning::MemoryUpdateStrategy;
use scirs2_core::ndarray::{Array1, Ix1};
fn quadratic_gradient(parameters: &Array1<f64>, target: &Array1<f64>) -> Array1<f64> {
parameters
.iter()
.zip(target.iter())
.map(|(&x, &t)| 2.0 * (x - t))
.collect()
}
fn quadratic_loss(parameters: &Array1<f64>, target: &Array1<f64>) -> f64 {
parameters
.iter()
.zip(target.iter())
.map(|(&x, &t)| (x - t) * (x - t))
.sum()
}
fn optimizer() -> LifelongOptimizer<f64, Ix1> {
LifelongOptimizer::new(LifelongStrategy::MemoryAugmented {
memory_size: 128,
update_strategy: MemoryUpdateStrategy::FIFO,
})
}
fn train(
optimizer: &mut LifelongOptimizer<f64, Ix1>,
task_id: &str,
target: &Array1<f64>,
steps: usize,
) {
for _ in 0..steps {
let parameters = optimizer
.task_parameters(task_id)
.expect("task must exist")
.clone();
let gradient = quadratic_gradient(¶meters, target);
let loss = quadratic_loss(¶meters, target);
optimizer
.update_current_task(&gradient, loss)
.expect("update must succeed");
}
}
#[test]
fn a_similar_task_is_warm_started_and_drops_its_loss_faster() {
let target_a = Array1::from_vec(vec![3.0, 3.0, 3.0, 3.0]);
let target_b = Array1::from_vec(vec![3.05, 3.05, 3.05, 3.05]);
let cold_start = Array1::zeros(4);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 4000);
let learned_a = opt.task_parameters("a").expect("task a").clone();
assert!(
quadratic_loss(&learned_a, &target_a) < 1.0,
"task a did not learn: {learned_a:?}"
);
let probe = quadratic_gradient(&cold_start, &target_b);
let outcome = opt
.start_task_with_probe("b".to_string(), cold_start.clone(), &probe)
.expect("start b");
assert_eq!(
outcome.source_task.as_deref(),
Some("a"),
"a task pulling the same way must transfer (similarity {})",
outcome.similarity
);
assert!(
outcome.similarity > 0.9,
"similarity {} is implausibly low for near-identical tasks",
outcome.similarity
);
assert!(outcome.transfer_weight > 0.9);
let warm_parameters = opt.task_parameters("b").expect("task b").clone();
let warm_initial_loss = quadratic_loss(&warm_parameters, &target_b);
let cold_initial_loss = quadratic_loss(&cold_start, &target_b);
assert!(
warm_initial_loss < cold_initial_loss * 0.1,
"warm start did not help: {warm_initial_loss} vs cold {cold_initial_loss}"
);
train(&mut opt, "b", &target_b, 20);
let warm_after = quadratic_loss(opt.task_parameters("b").expect("task b"), &target_b);
let mut cold = optimizer();
cold.start_task("b".to_string(), cold_start.clone())
.expect("cold start b");
train(&mut cold, "b", &target_b, 20);
let cold_after = quadratic_loss(cold.task_parameters("b").expect("cold task b"), &target_b);
assert!(
warm_after < cold_after,
"warm start ({warm_after}) must beat cold start ({cold_after}) after 20 steps"
);
}
#[test]
fn a_dissimilar_task_is_not_warm_started() {
let target_a = Array1::from_vec(vec![3.0, 3.0, 3.0, 3.0]);
let target_c = Array1::from_vec(vec![-3.0, -3.0, -3.0, -3.0]);
let cold_start = Array1::zeros(4);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 500);
let probe = quadratic_gradient(&cold_start, &target_c);
let outcome = opt
.start_task_with_probe("c".to_string(), cold_start.clone(), &probe)
.expect("start c");
assert_eq!(
outcome.source_task, None,
"an opposing task must not transfer (similarity {})",
outcome.similarity
);
assert!(
outcome.similarity < 0.0,
"opposing gradients must score below zero, got {}",
outcome.similarity
);
assert_eq!(outcome.transfer_weight, 0.0);
assert_eq!(
opt.task_parameters("c").expect("task c"),
&cold_start,
"a task that did not transfer must start exactly where the caller put it"
);
assert!(opt.task_dependencies("c").is_empty());
}
#[test]
fn transfer_weights_and_dependencies_are_recorded() {
let target_a = Array1::from_vec(vec![2.0, -2.0]);
let target_b = Array1::from_vec(vec![2.1, -2.1]);
let cold_start = Array1::zeros(2);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 500);
let probe = quadratic_gradient(&cold_start, &target_b);
opt.start_task_with_probe("b".to_string(), cold_start, &probe)
.expect("start b");
assert_eq!(opt.task_dependencies("b").to_vec(), vec!["a".to_string()]);
assert!(opt.transfer_weight("a", "b") > 0.9);
assert_eq!(
opt.transfer_weight("b", "a"),
0.0,
"transfer is directed: b was started from a, not the other way round"
);
assert!(
opt.compute_task_similarity("a", "b") > 0.9,
"the similarity matrix must be populated"
);
assert!(
opt.get_lifelong_stats().transfer_efficiency > 0.9,
"transfer efficiency must reflect the transfers that happened"
);
assert!(opt.task_embedding("a").is_some());
assert!(opt.task_embedding("b").is_some());
}
#[test]
fn similar_tasks_cluster_together() {
let cold_start = Array1::zeros(3);
let target_a = Array1::from_vec(vec![1.0, 1.0, 1.0]);
let target_b = Array1::from_vec(vec![1.2, 1.1, 1.05]);
let target_c = Array1::from_vec(vec![-1.0, -1.0, -1.0]);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 200);
opt.start_task("b".to_string(), cold_start.clone())
.expect("start b");
train(&mut opt, "b", &target_b, 200);
opt.start_task("c".to_string(), cold_start)
.expect("start c");
train(&mut opt, "c", &target_c, 200);
let clusters = opt.task_clusters().to_vec();
assert_eq!(
clusters,
vec![
vec!["a".to_string(), "b".to_string()],
vec!["c".to_string()]
],
"clusters: {clusters:?}, sim(a,b) = {}, sim(a,c) = {}",
opt.compute_task_similarity("a", "b"),
opt.compute_task_similarity("a", "c")
);
}
#[test]
fn the_threshold_controls_clustering_and_transfer() {
let cold_start = Array1::zeros(6);
let target_a = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0]);
let target_b = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0, -1.0]);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 200);
opt.start_task("b".to_string(), cold_start.clone())
.expect("start b");
train(&mut opt, "b", &target_b, 200);
let similarity = opt.compute_task_similarity("a", "b");
assert!(
(0.5..0.9).contains(&similarity),
"five agreeing coordinates out of six should score in [0.5, 0.9), got {similarity}"
);
assert_eq!(
opt.task_clusters().len(),
1,
"similar tasks share a cluster"
);
opt.set_transfer_threshold(0.9).expect("valid threshold");
assert_eq!(
opt.task_clusters().len(),
2,
"a threshold above the measured similarity must separate the tasks"
);
let target_d = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, -1.0, 1.0]);
let probe = quadratic_gradient(&cold_start, &target_d);
let outcome = opt
.start_task_with_probe("d".to_string(), cold_start, &probe)
.expect("start d");
assert!(
(0.5..0.9).contains(&outcome.similarity),
"probe similarity {} is outside the band this test needs",
outcome.similarity
);
assert_eq!(
outcome.source_task, None,
"similarity {} cleared a threshold of 0.9",
outcome.similarity
);
assert!(opt.set_transfer_threshold(2.0).is_err());
assert!(opt.set_transfer_threshold(f64::NAN).is_err());
}
#[test]
fn a_mismatched_probe_is_reported() {
let mut opt = optimizer();
let error = opt
.start_task_with_probe(
"a".to_string(),
Array1::zeros(3),
&Array1::from_vec(vec![1.0, 2.0]),
)
.expect_err("a shape mismatch must be reported");
assert!(
matches!(error, OptimError::DimensionMismatch(_)),
"{error:?}"
);
}
#[test]
fn the_first_task_has_nothing_to_transfer_from() {
let mut opt = optimizer();
let outcome = opt
.start_task_with_probe(
"a".to_string(),
Array1::zeros(2),
&Array1::from_vec(vec![1.0, -1.0]),
)
.expect("start a");
assert_eq!(outcome.source_task, None);
assert_eq!(outcome.similarity, 0.0);
assert_eq!(opt.mean_transfer_weight(), 0.0);
assert_eq!(opt.task_clusters().to_vec(), vec![vec!["a".to_string()]]);
}
#[test]
fn embeddings_have_a_fixed_width_regardless_of_parameter_count() {
let mut small = TaskStatistics::default();
small.observe(&[1.0, -1.0, 1.0]);
let mut large = TaskStatistics::default();
large.observe(&vec![0.5; 500]);
assert_eq!(small.embedding(64).len(), 64);
assert_eq!(large.embedding(64).len(), 64);
let similarity = cosine_similarity(&small.embedding(64), &large.embedding(64));
assert!(
(-1.0..=1.0).contains(&similarity),
"similarity {similarity} is not a cosine"
);
}
#[test]
fn the_embedding_captures_gradient_direction() {
let mut forward = TaskStatistics::default();
let mut backward = TaskStatistics::default();
for step in 0..10 {
let scale = 1.0 + step as f64;
forward.observe(&[scale, 2.0 * scale, -scale]);
backward.observe(&[-scale, -2.0 * scale, scale]);
}
let same = cosine_similarity(&forward.embedding(32), &forward.embedding(32));
let opposite = cosine_similarity(&forward.embedding(32), &backward.embedding(32));
assert!((same - 1.0).abs() < 1e-9, "same direction scored {same}");
assert!(
(opposite + 1.0).abs() < 1e-9,
"opposite direction scored {opposite}"
);
}
#[test]
fn a_directionless_task_has_a_zero_embedding() {
let mut statistics = TaskStatistics::default();
statistics.observe(&[1.0, 1.0]);
statistics.observe(&[-1.0, -1.0]);
let embedding = statistics.embedding(8);
assert!(embedding.iter().all(|&value| value == 0.0));
assert_eq!(cosine_similarity(&embedding, &embedding), 0.0);
}
#[test]
fn forgetting_is_measured_from_re_evaluated_tasks() {
let cold_start = Array1::zeros(2);
let target_a = Array1::from_vec(vec![1.0, 1.0]);
let mut opt = optimizer();
opt.start_task("a".to_string(), cold_start.clone())
.expect("start a");
train(&mut opt, "a", &target_a, 50);
assert_eq!(
opt.get_lifelong_stats().catastrophic_forgetting,
0.0,
"forgetting is not observable while the task is still current"
);
opt.start_task("b".to_string(), cold_start)
.expect("start b");
assert_eq!(
opt.get_lifelong_stats().catastrophic_forgetting,
0.0,
"switching away does not by itself demonstrate forgetting"
);
let reference = opt
.task_performance
.get("a")
.and_then(|history| history.last())
.copied()
.expect("task a recorded losses");
opt.record_task_performance("a", reference + 0.25)
.expect("recording a re-evaluation must succeed");
let forgetting = opt.get_lifelong_stats().catastrophic_forgetting;
assert!(
(forgetting - 0.25).abs() < 1e-9,
"forgetting was not measured from the re-evaluation: {forgetting}"
);
assert!(opt.record_task_performance("b", 1.0).is_err());
assert!(opt.record_task_performance("nope", 1.0).is_err());
}
#[test]
fn statistics_are_running_means() {
let mut statistics = TaskStatistics::default();
statistics.observe(&[2.0, 0.0]);
statistics.observe(&[4.0, 0.0]);
assert_eq!(statistics.observations(), 2);
assert!((statistics.mean_gradient()[0] - 3.0).abs() < 1e-12);
assert!((statistics.mean_squared_gradient()[0] - 10.0).abs() < 1e-12);
}
}