use super::meta_learning::MetaExperience;
use scirs2_core::numeric::Float;
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransferStrategy {
ParameterTransfer,
FeatureTransfer,
InstanceTransfer,
RelationalTransfer,
MetaTransfer,
}
impl TransferStrategy {
pub fn is_supported(self) -> bool {
matches!(self, TransferStrategy::InstanceTransfer)
}
}
#[derive(Debug, Clone)]
pub struct DomainAdaptation<A: Float + Send + Sync> {
source_characteristics: Vec<A>,
target_characteristics: Vec<A>,
adaptation_weights: Vec<A>,
domain_similarity: Option<A>,
}
impl<A: Float + Send + Sync> Default for DomainAdaptation<A> {
fn default() -> Self {
Self {
source_characteristics: Vec::new(),
target_characteristics: Vec::new(),
adaptation_weights: Vec::new(),
domain_similarity: None,
}
}
}
impl<A: Float + Send + Sync> DomainAdaptation<A> {
pub fn set_source_characteristics(&mut self, characteristics: Vec<A>) {
self.source_characteristics = characteristics;
self.recompute();
}
pub fn set_target_characteristics(&mut self, characteristics: Vec<A>) {
self.target_characteristics = characteristics;
self.recompute();
}
pub fn domain_similarity(&self) -> Option<A> {
self.domain_similarity
}
pub fn adaptation_weights(&self) -> &[A] {
&self.adaptation_weights
}
fn recompute(&mut self) {
let shared = self
.source_characteristics
.len()
.min(self.target_characteristics.len());
if shared == 0 {
self.domain_similarity = None;
self.adaptation_weights.clear();
return;
}
let mut dot = A::zero();
let mut source_norm = A::zero();
let mut target_norm = A::zero();
self.adaptation_weights.clear();
for index in 0..shared {
let source = self.source_characteristics[index];
let target = self.target_characteristics[index];
dot = dot + source * target;
source_norm = source_norm + source * source;
target_norm = target_norm + target * target;
self.adaptation_weights.push(if source == A::zero() {
A::one()
} else {
target / source
});
}
self.domain_similarity = if source_norm > A::zero() && target_norm > A::zero() {
Some(dot / (source_norm.sqrt() * target_norm.sqrt()))
} else {
None
};
}
}
#[derive(Debug, Clone)]
pub struct TransferMetrics<A: Float + Send + Sync> {
pub success_rate: Option<A>,
pub improvement: Option<A>,
pub efficiency: Option<A>,
pub negative_transfer_count: usize,
pub evaluated_transfers: usize,
}
impl<A: Float + Send + Sync> Default for TransferMetrics<A> {
fn default() -> Self {
Self {
success_rate: None,
improvement: None,
efficiency: None,
negative_transfer_count: 0,
evaluated_transfers: 0,
}
}
}
pub const MIN_TRANSFER_SIMILARITY: f64 = 0.5;
#[derive(Debug, Clone)]
pub struct TransferLearning<A: Float + Send + Sync> {
source_experiences: HashMap<String, Vec<MetaExperience<A>>>,
transfer_strategies: Vec<TransferStrategy>,
domain_adaptation: DomainAdaptation<A>,
transfer_metrics: TransferMetrics<A>,
improvement_total: A,
successful_transfers: usize,
}
impl<A: Float + Send + Sync + Clone> Default for TransferLearning<A> {
fn default() -> Self {
Self::new()
}
}
impl<A: Float + Send + Sync + Clone> TransferLearning<A> {
pub fn new() -> Self {
Self {
source_experiences: HashMap::new(),
transfer_strategies: vec![TransferStrategy::InstanceTransfer],
domain_adaptation: DomainAdaptation::default(),
transfer_metrics: TransferMetrics::default(),
improvement_total: A::zero(),
successful_transfers: 0,
}
}
pub fn set_strategies(&mut self, strategies: Vec<TransferStrategy>) -> Result<(), String> {
if let Some(unsupported) = strategies.iter().find(|s| !s.is_supported()) {
return Err(format!(
"transfer strategy {unsupported:?} is not implemented: the meta-learner \
stores experiences, not model parameters, feature extractors or \
relational structure; only InstanceTransfer is supported"
));
}
if strategies.is_empty() {
return Err("at least one transfer strategy is required".to_string());
}
self.transfer_strategies = strategies;
Ok(())
}
pub fn strategies(&self) -> &[TransferStrategy] {
&self.transfer_strategies
}
pub fn register_source(
&mut self,
source_id: String,
experiences: Vec<MetaExperience<A>>,
characteristics: Vec<A>,
) {
self.domain_adaptation
.set_source_characteristics(characteristics);
self.source_experiences.insert(source_id, experiences);
}
pub fn source_domain_count(&self) -> usize {
self.source_experiences.len()
}
pub fn metrics(&self) -> &TransferMetrics<A> {
&self.transfer_metrics
}
pub fn domain_similarity(&self) -> Option<A> {
self.domain_adaptation.domain_similarity()
}
pub fn select_transfer_batch(
&mut self,
target_characteristics: Vec<A>,
limit: usize,
) -> Vec<MetaExperience<A>> {
self.domain_adaptation
.set_target_characteristics(target_characteristics);
let available: usize = self.source_experiences.values().map(Vec::len).sum();
if available == 0 || limit == 0 {
return Vec::new();
}
let Some(similarity) = self.domain_adaptation.domain_similarity() else {
return Vec::new();
};
let threshold = A::from(MIN_TRANSFER_SIMILARITY).unwrap_or_else(A::zero);
if similarity < threshold {
self.record_efficiency(0, available);
return Vec::new();
}
let mut candidates: Vec<MetaExperience<A>> = self
.source_experiences
.values()
.flat_map(|batch| batch.iter().cloned())
.collect();
candidates.sort_by(|a, b| crate::utils::total_order(&b.priority, &a.priority));
candidates.truncate(limit);
for experience in &mut candidates {
experience.priority = experience.priority * similarity;
}
self.record_efficiency(candidates.len(), available);
candidates
}
pub fn record_transfer_outcome(&mut self, reward_before: A, reward_after: A) {
let improvement = reward_after - reward_before;
self.transfer_metrics.evaluated_transfers += 1;
self.improvement_total = self.improvement_total + improvement;
if improvement > A::zero() {
self.successful_transfers += 1;
} else if improvement < A::zero() {
self.transfer_metrics.negative_transfer_count += 1;
}
if let Some(evaluated) = A::from(self.transfer_metrics.evaluated_transfers) {
if evaluated > A::zero() {
self.transfer_metrics.improvement = Some(self.improvement_total / evaluated);
self.transfer_metrics.success_rate =
A::from(self.successful_transfers).map(|successes| successes / evaluated);
}
}
}
fn record_efficiency(&mut self, retained: usize, available: usize) {
if available == 0 {
self.transfer_metrics.efficiency = None;
return;
}
let (Some(retained), Some(available)) = (A::from(retained), A::from(available)) else {
return;
};
self.transfer_metrics.efficiency = Some(retained / available);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::streaming::adaptive_streaming::meta_learning::{
EpisodeContext, EpisodeOutcome, MetaAction, MetaExperience, MetaState,
};
use crate::streaming::adaptive_streaming::optimizer::AdaptationType;
use std::time::{Duration, Instant};
fn experience(priority: f64, reward: f64) -> MetaExperience<f64> {
MetaExperience {
id: 1,
state: MetaState {
performance_metrics: vec![reward],
resource_state: vec![1.0],
drift_indicators: vec![0.0],
adaptation_history: 0,
timestamp: Instant::now(),
},
action: MetaAction {
adaptation_magnitudes: vec![0.1],
adaptation_types: vec![AdaptationType::LearningRate],
learning_rate_change: 0.1,
buffer_size_change: 0.0,
timestamp: Instant::now(),
},
reward,
next_state: None,
timestamp: Instant::now(),
episode_context: EpisodeContext {
episode_id: 0,
start_time: Instant::now(),
duration: Duration::ZERO,
initial_performance: 0.0,
final_performance: reward,
adaptation_count: 1,
outcome: EpisodeOutcome::Neutral,
},
priority,
replay_count: 0,
}
}
#[test]
fn fresh_metrics_are_absent_not_fabricated() {
let transfer = TransferLearning::<f64>::new();
let metrics = transfer.metrics();
assert!(metrics.success_rate.is_none());
assert!(metrics.improvement.is_none());
assert!(metrics.efficiency.is_none());
assert_eq!(metrics.evaluated_transfers, 0);
assert!(
transfer.domain_similarity().is_none(),
"no domain has been described yet, so there is no similarity to report"
);
}
#[test]
fn similar_domains_transfer_with_similarity_scaled_priorities() {
let mut transfer = TransferLearning::<f64>::new();
transfer.register_source(
"source".to_string(),
vec![experience(1.0, 1.0), experience(0.2, 0.0)],
vec![1.0, 1.0],
);
let batch = transfer.select_transfer_batch(vec![1.0, 1.0], 2);
assert_eq!(batch.len(), 2, "identical domains must transfer everything");
let similarity = transfer.domain_similarity().expect("similarity");
assert!((similarity - 1.0).abs() < 1e-9, "similarity = {similarity}");
assert!((batch[0].priority - 1.0).abs() < 1e-9);
assert_eq!(transfer.metrics().efficiency, Some(1.0));
}
#[test]
fn dissimilar_domains_transfer_nothing() {
let mut transfer = TransferLearning::<f64>::new();
transfer.register_source(
"source".to_string(),
vec![experience(1.0, 1.0)],
vec![1.0, 0.0],
);
let batch = transfer.select_transfer_batch(vec![0.0, 1.0], 4);
assert!(batch.is_empty(), "orthogonal domains must not transfer");
assert_eq!(transfer.metrics().efficiency, Some(0.0));
}
#[test]
fn transfer_outcomes_are_measured() {
let mut transfer = TransferLearning::<f64>::new();
transfer.record_transfer_outcome(1.0, 2.0); transfer.record_transfer_outcome(1.0, 0.0);
let metrics = transfer.metrics();
assert_eq!(metrics.evaluated_transfers, 2);
assert_eq!(metrics.negative_transfer_count, 1);
assert_eq!(metrics.success_rate, Some(0.5));
assert_eq!(metrics.improvement, Some(0.0));
}
#[test]
fn unsupported_strategies_are_refused() {
let mut transfer = TransferLearning::<f64>::new();
let err = transfer
.set_strategies(vec![TransferStrategy::ParameterTransfer])
.expect_err("ParameterTransfer must be refused");
assert!(err.contains("ParameterTransfer"), "{err}");
transfer
.set_strategies(vec![TransferStrategy::InstanceTransfer])
.expect("InstanceTransfer is supported");
assert_eq!(transfer.strategies(), [TransferStrategy::InstanceTransfer]);
}
}