use crate::layers::Layer;
use crate::weight_init::{InitStrategy, WeightInitializer};
use crate::NeuralResult;
use scirs2_core::ndarray::{Array1, Array2, Array3};
use scirs2_core::random::ChaCha8Rng;
use scirs2_core::random::SeedableRng;
use sklears_core::types::FloatBounds;
use std::collections::{HashMap, HashSet};
use std::marker::PhantomData;
#[derive(Debug, Clone)]
pub struct TransferConfig<T: FloatBounds> {
pub frozen_layers: HashSet<String>,
pub layer_learning_rates: HashMap<String, T>,
pub fine_tuning_strategy: FineTuningStrategy,
pub unfreeze_schedule: Option<UnfreezeSchedule>,
pub domain_adaptation: Option<DomainAdaptationConfig<T>>,
pub discriminative_lr: bool,
pub base_learning_rate: T,
}
impl<T: FloatBounds> Default for TransferConfig<T> {
fn default() -> Self {
Self {
frozen_layers: HashSet::new(),
layer_learning_rates: HashMap::new(),
fine_tuning_strategy: FineTuningStrategy::FineTuneAll,
unfreeze_schedule: None,
domain_adaptation: None,
discriminative_lr: false,
base_learning_rate: T::from(1e-3).unwrap_or_else(|| T::zero()),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum FineTuningStrategy {
FeatureExtraction,
FineTuneAll,
FineTuneTop {
num_layers: usize,
},
GradualUnfreeze,
LayerWiseAdaptive,
TaskSpecific,
}
#[derive(Debug, Clone)]
pub struct UnfreezeSchedule {
pub epochs_per_step: usize,
pub layers_per_step: usize,
pub unfreeze_direction: UnfreezeDirection,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum UnfreezeDirection {
TopToBottom,
BottomToTop,
MiddleOut,
}
#[derive(Debug, Clone)]
pub struct DomainAdaptationConfig<T: FloatBounds> {
pub technique: DomainAdaptationTechnique,
pub adaptation_weight: T,
pub adaptation_iterations: usize,
pub adversarial_training: bool,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum DomainAdaptationTechnique {
MMD,
DANN,
CORAL,
AdaBN,
}
#[derive(Debug, Clone)]
pub struct TransferLearningManager<T: FloatBounds> {
config: TransferConfig<T>,
current_epoch: usize,
layer_freeze_status: HashMap<String, bool>,
layer_lr_multipliers: HashMap<String, T>,
training_history: Vec<TransferMetrics<T>>,
}
impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> TransferLearningManager<T> {
pub fn new(config: TransferConfig<T>) -> Self {
let mut layer_freeze_status = HashMap::new();
let mut layer_lr_multipliers = HashMap::new();
for layer_name in &config.frozen_layers {
layer_freeze_status.insert(layer_name.clone(), true);
}
for (layer_name, lr_mult) in &config.layer_learning_rates {
layer_lr_multipliers.insert(layer_name.clone(), *lr_mult);
}
Self {
config,
current_epoch: 0,
layer_freeze_status,
layer_lr_multipliers,
training_history: Vec::new(),
}
}
pub fn freeze_layers(&mut self, layer_names: &[String]) -> NeuralResult<()> {
for layer_name in layer_names {
self.layer_freeze_status.insert(layer_name.clone(), true);
}
Ok(())
}
pub fn unfreeze_layers(&mut self, layer_names: &[String]) -> NeuralResult<()> {
for layer_name in layer_names {
self.layer_freeze_status.insert(layer_name.clone(), false);
}
Ok(())
}
pub fn is_layer_frozen(&self, layer_name: &str) -> bool {
self.layer_freeze_status
.get(layer_name)
.copied()
.unwrap_or(false)
}
pub fn get_layer_lr_multiplier(&self, layer_name: &str) -> T {
self.layer_lr_multipliers
.get(layer_name)
.copied()
.unwrap_or(T::one())
}
pub fn update_epoch(&mut self, epoch: usize) -> NeuralResult<()> {
self.current_epoch = epoch;
if let Some(schedule) = self.config.unfreeze_schedule.clone() {
self.apply_unfreeze_schedule(&schedule)?;
}
Ok(())
}
fn apply_unfreeze_schedule(&mut self, schedule: &UnfreezeSchedule) -> NeuralResult<()> {
if self.current_epoch.is_multiple_of(schedule.epochs_per_step) && self.current_epoch > 0 {
let step = self.current_epoch / schedule.epochs_per_step;
let layers_to_unfreeze = self.select_layers_for_unfreezing(schedule, step)?;
self.unfreeze_layers(&layers_to_unfreeze)?;
}
Ok(())
}
fn select_layers_for_unfreezing(
&self,
schedule: &UnfreezeSchedule,
step: usize,
) -> NeuralResult<Vec<String>> {
let frozen_layers: Vec<String> = self
.layer_freeze_status
.iter()
.filter(|(_, &is_frozen)| is_frozen)
.map(|(name, _)| name.clone())
.collect();
let start_idx = step * schedule.layers_per_step;
let end_idx = ((step + 1) * schedule.layers_per_step).min(frozen_layers.len());
if start_idx >= frozen_layers.len() {
return Ok(Vec::new());
}
let selected_layers = match schedule.unfreeze_direction {
UnfreezeDirection::TopToBottom => frozen_layers[start_idx..end_idx].to_vec(),
UnfreezeDirection::BottomToTop => {
let mut layers = frozen_layers.clone();
layers.reverse();
layers[start_idx..end_idx].to_vec()
}
UnfreezeDirection::MiddleOut => {
let middle = frozen_layers.len() / 2;
let mut selected = Vec::new();
for i in 0..schedule.layers_per_step {
if step * schedule.layers_per_step + i >= frozen_layers.len() {
break;
}
let offset = i / 2;
if i % 2 == 0 {
if middle + offset < frozen_layers.len() {
selected.push(frozen_layers[middle + offset].clone());
}
} else {
if offset < middle {
selected.push(frozen_layers[middle - offset - 1].clone());
}
}
}
selected
}
};
Ok(selected_layers)
}
pub fn apply_discriminative_learning_rates(
&mut self,
layer_names: &[String],
) -> NeuralResult<()> {
if !self.config.discriminative_lr {
return Ok(());
}
let num_layers = layer_names.len();
for (i, layer_name) in layer_names.iter().enumerate() {
let layer_depth = (num_layers - i - 1) as f64;
let lr_multiplier =
T::from(0.1_f64.powf(layer_depth / num_layers as f64)).unwrap_or_else(|| T::zero());
self.layer_lr_multipliers
.insert(layer_name.clone(), lr_multiplier);
}
Ok(())
}
pub fn record_metrics(&mut self, metrics: TransferMetrics<T>) {
self.training_history.push(metrics);
}
pub fn get_training_history(&self) -> &[TransferMetrics<T>] {
&self.training_history
}
pub fn calculate_transfer_effectiveness(&self) -> Option<T> {
if self.training_history.len() < 2 {
return None;
}
let initial_loss = self.training_history[0].validation_loss;
let final_loss = self
.training_history
.last()
.expect("empty collection")
.validation_loss;
Some((initial_loss - final_loss) / initial_loss)
}
}
#[derive(Debug, Clone)]
pub struct TransferMetrics<T: FloatBounds> {
pub epoch: usize,
pub training_loss: T,
pub validation_loss: T,
pub frozen_layer_count: usize,
pub lr_stats: LearningRateStats<T>,
pub domain_adaptation_loss: Option<T>,
}
#[derive(Debug, Clone)]
pub struct LearningRateStats<T: FloatBounds> {
pub mean_lr: T,
pub max_lr: T,
pub min_lr: T,
pub lr_variance: T,
}
pub struct ModelAdapter<T: FloatBounds> {
layer_replacements: HashMap<String, Box<dyn Layer<T>>>,
init_strategy: InitStrategy,
}
impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> ModelAdapter<T> {
pub fn new(init_strategy: InitStrategy) -> Self {
Self {
layer_replacements: HashMap::new(),
init_strategy,
}
}
pub fn replace_layer(&mut self, layer_name: String, new_layer: Box<dyn Layer<T>>) {
self.layer_replacements.insert(layer_name, new_layer);
}
pub fn replace_classifier(
&mut self,
num_classes: usize,
hidden_size: usize,
_layer_name: String,
) -> NeuralResult<()> {
let mut rng = ChaCha8Rng::seed_from_u64(42);
let initializer: WeightInitializer<T> = WeightInitializer::new(self.init_strategy);
let _weights = initializer.initialize_2d(&mut rng, (hidden_size, num_classes))?;
let _bias: Array1<T> = Array1::zeros(num_classes);
Ok(())
}
pub fn apply_replacements<M>(&self, model: &mut M) -> NeuralResult<()>
where
M: HasReplaceableLayer<T>,
{
for (layer_name, replacement_layer) in &self.layer_replacements {
model.replace_layer(layer_name, replacement_layer.as_ref())?;
}
Ok(())
}
}
impl<T: FloatBounds> std::fmt::Debug for ModelAdapter<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModelAdapter")
.field(
"layer_replacements",
&format!("{} layers", self.layer_replacements.len()),
)
.field("init_strategy", &self.init_strategy)
.finish()
}
}
pub trait HasReplaceableLayer<T: FloatBounds> {
fn replace_layer(&mut self, layer_name: &str, new_layer: &dyn Layer<T>) -> NeuralResult<()>;
fn get_layer_names(&self) -> Vec<String>;
}
#[derive(Debug, Clone)]
pub struct FeatureExtractor<T: FloatBounds> {
extract_layer: String,
global_pooling: bool,
pooling_strategy: PoolingStrategy,
_phantom: PhantomData<T>,
}
impl<T: FloatBounds> FeatureExtractor<T> {
pub fn new(
extract_layer: String,
global_pooling: bool,
pooling_strategy: PoolingStrategy,
) -> Self {
Self {
extract_layer,
global_pooling,
pooling_strategy,
_phantom: PhantomData,
}
}
pub fn extract_features<M>(&self, model: &M, input: &Array3<T>) -> NeuralResult<Array2<T>>
where
M: HasFeatureExtraction<T>,
{
let features = model.extract_features_at_layer(&self.extract_layer, input)?;
if self.global_pooling {
self.apply_global_pooling(&features)
} else {
let (batch_size, _, _) = features.dim();
let flattened_size = features.len() / batch_size;
Ok(features
.into_shape_with_order((batch_size, flattened_size))
.expect("array shape error"))
}
}
fn apply_global_pooling(&self, features: &Array3<T>) -> NeuralResult<Array2<T>> {
let (batch_size, _, num_features) = features.dim();
let mut pooled = Array2::zeros((batch_size, num_features));
for b in 0..batch_size {
for f in 0..num_features {
let feature_map = features.slice(scirs2_core::ndarray::s![b, .., f]);
let pooled_value = match self.pooling_strategy {
PoolingStrategy::Mean => feature_map.mean().unwrap_or(T::zero()),
PoolingStrategy::Max => feature_map.fold(T::zero(), |acc, &x| acc.max(x)),
PoolingStrategy::Min => feature_map.fold(T::zero(), |acc, &x| acc.min(x)),
};
pooled[[b, f]] = pooled_value;
}
}
Ok(pooled)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum PoolingStrategy {
Mean,
Max,
Min,
}
pub trait HasFeatureExtraction<T: FloatBounds> {
fn extract_features_at_layer(
&self,
layer_name: &str,
input: &Array3<T>,
) -> NeuralResult<Array3<T>>;
}
#[derive(Debug, Clone)]
pub struct DomainAdapter<T: FloatBounds> {
config: DomainAdaptationConfig<T>,
source_stats: Option<DomainStatistics<T>>,
target_stats: Option<DomainStatistics<T>>,
}
impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> DomainAdapter<T> {
pub fn new(config: DomainAdaptationConfig<T>) -> Self {
Self {
config,
source_stats: None,
target_stats: None,
}
}
pub fn compute_domain_statistics(
&mut self,
source_data: &Array3<T>,
target_data: &Array3<T>,
) -> NeuralResult<()> {
self.source_stats = Some(self.compute_statistics(source_data)?);
self.target_stats = Some(self.compute_statistics(target_data)?);
Ok(())
}
pub fn adapt_features(&self, features: &Array3<T>, is_source: bool) -> NeuralResult<Array3<T>> {
match self.config.technique {
DomainAdaptationTechnique::CORAL => self.apply_coral_adaptation(features, is_source),
DomainAdaptationTechnique::AdaBN => self.apply_adaptive_batch_norm(features, is_source),
_ => {
Ok(features.clone())
}
}
}
fn compute_statistics(&self, data: &Array3<T>) -> NeuralResult<DomainStatistics<T>> {
let (batch_size, seq_len, feature_dim) = data.dim();
let total_samples = batch_size * seq_len;
let mut mean = Array1::zeros(feature_dim);
for b in 0..batch_size {
for s in 0..seq_len {
for f in 0..feature_dim {
mean[f] += data[[b, s, f]];
}
}
}
mean /= T::from(total_samples).unwrap_or_else(|| T::zero());
let mut covariance = Array2::zeros((feature_dim, feature_dim));
for b in 0..batch_size {
for s in 0..seq_len {
for i in 0..feature_dim {
for j in 0..feature_dim {
let diff_i = data[[b, s, i]] - mean[i];
let diff_j = data[[b, s, j]] - mean[j];
covariance[[i, j]] += diff_i * diff_j;
}
}
}
}
covariance /= T::from(total_samples - 1).unwrap_or_else(|| T::zero());
Ok(DomainStatistics { mean, covariance })
}
fn apply_coral_adaptation(
&self,
features: &Array3<T>,
is_source: bool,
) -> NeuralResult<Array3<T>> {
let stats = if is_source {
&self.source_stats
} else {
&self.target_stats
};
if let Some(domain_stats) = stats {
let (batch_size, seq_len, feature_dim) = features.dim();
let mut adapted = features.clone();
for b in 0..batch_size {
for s in 0..seq_len {
for f in 0..feature_dim {
adapted[[b, s, f]] -= domain_stats.mean[f];
}
}
}
Ok(adapted)
} else {
Ok(features.clone())
}
}
fn apply_adaptive_batch_norm(
&self,
features: &Array3<T>,
_is_source: bool,
) -> NeuralResult<Array3<T>> {
let (batch_size, seq_len, feature_dim) = features.dim();
let mut normalized = Array3::zeros((batch_size, seq_len, feature_dim));
for f in 0..feature_dim {
let mut sum = T::zero();
let mut count = 0;
for b in 0..batch_size {
for s in 0..seq_len {
sum += features[[b, s, f]];
count += 1;
}
}
let mean = sum / T::from(count).unwrap_or_else(|| T::zero());
let mut variance_sum = T::zero();
for b in 0..batch_size {
for s in 0..seq_len {
let diff = features[[b, s, f]] - mean;
variance_sum += diff * diff;
}
}
let variance = variance_sum / T::from(count).unwrap_or_else(|| T::zero());
let std = (variance + T::from(1e-5).unwrap_or_else(|| T::zero())).sqrt();
for b in 0..batch_size {
for s in 0..seq_len {
normalized[[b, s, f]] = (features[[b, s, f]] - mean) / std;
}
}
}
Ok(normalized)
}
}
#[derive(Debug, Clone)]
pub struct DomainStatistics<T: FloatBounds> {
pub mean: Array1<T>,
pub covariance: Array2<T>,
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::essentials::Normal;
use scirs2_core::ndarray::Array3;
use scirs2_core::random::thread_rng;
#[test]
fn test_transfer_learning_manager_creation() {
let config = TransferConfig::default();
let manager = TransferLearningManager::<f64>::new(config);
assert_eq!(manager.current_epoch, 0);
}
#[test]
fn test_layer_freezing() {
let mut config = TransferConfig::<f64>::default();
config.frozen_layers.insert("layer1".to_string());
let mut manager = TransferLearningManager::new(config);
assert!(manager.is_layer_frozen("layer1"));
assert!(!manager.is_layer_frozen("layer2"));
manager
.unfreeze_layers(&["layer1".to_string()])
.expect("operation should succeed");
assert!(!manager.is_layer_frozen("layer1"));
}
#[test]
fn test_learning_rate_multipliers() {
let mut config = TransferConfig::<f64>::default();
config
.layer_learning_rates
.insert("layer1".to_string(), 0.5);
let manager = TransferLearningManager::new(config);
assert_eq!(manager.get_layer_lr_multiplier("layer1"), 0.5);
assert_eq!(manager.get_layer_lr_multiplier("layer2"), 1.0);
}
#[test]
fn test_model_adapter_creation() {
let adapter = ModelAdapter::<f64>::new(InitStrategy::XavierUniform);
assert_eq!(adapter.layer_replacements.len(), 0);
}
#[test]
fn test_feature_extractor() {
let extractor =
FeatureExtractor::<f64>::new("conv_layer".to_string(), true, PoolingStrategy::Mean);
assert_eq!(extractor.extract_layer, "conv_layer");
assert!(extractor.global_pooling);
}
#[test]
fn test_domain_adapter() {
let config = DomainAdaptationConfig {
technique: DomainAdaptationTechnique::CORAL,
adaptation_weight: 0.1,
adaptation_iterations: 100,
adversarial_training: false,
};
let adapter = DomainAdapter::<f64>::new(config);
assert!(adapter.source_stats.is_none());
assert!(adapter.target_stats.is_none());
}
#[test]
fn test_domain_statistics_computation() {
let config = DomainAdaptationConfig {
technique: DomainAdaptationTechnique::CORAL,
adaptation_weight: 0.1,
adaptation_iterations: 100,
adversarial_training: false,
};
let mut adapter = DomainAdapter::new(config);
let source_data = Array3::from_shape_fn((10, 5, 8), |_| {
let mut rng = thread_rng();
rng.sample(Normal::new(0.0, 1.0).expect("construction should succeed"))
});
let target_data = Array3::from_shape_fn((10, 5, 8), |_| {
let mut rng = thread_rng();
rng.sample(Normal::new(0.0, 1.0).expect("construction should succeed"))
});
let result = adapter.compute_domain_statistics(&source_data, &target_data);
assert!(result.is_ok());
assert!(adapter.source_stats.is_some());
assert!(adapter.target_stats.is_some());
}
#[test]
fn test_unfreeze_schedule() {
let schedule = UnfreezeSchedule {
epochs_per_step: 5,
layers_per_step: 2,
unfreeze_direction: UnfreezeDirection::TopToBottom,
};
let mut config = TransferConfig::<f64> {
unfreeze_schedule: Some(schedule),
..Default::default()
};
for i in 0..6 {
config.frozen_layers.insert(format!("layer_{}", i));
}
let mut manager = TransferLearningManager::new(config);
manager.update_epoch(5).expect("operation should succeed");
let frozen_count = manager.layer_freeze_status.values().filter(|&&v| v).count();
assert!(frozen_count < 6);
}
#[test]
fn test_transfer_effectiveness_calculation() {
let config = TransferConfig::default();
let mut manager = TransferLearningManager::new(config);
manager.record_metrics(TransferMetrics {
epoch: 0,
training_loss: 1.0,
validation_loss: 1.0,
frozen_layer_count: 5,
lr_stats: LearningRateStats {
mean_lr: 0.001,
max_lr: 0.001,
min_lr: 0.001,
lr_variance: 0.0,
},
domain_adaptation_loss: None,
});
manager.record_metrics(TransferMetrics {
epoch: 10,
training_loss: 0.5,
validation_loss: 0.6,
frozen_layer_count: 3,
lr_stats: LearningRateStats {
mean_lr: 0.001,
max_lr: 0.001,
min_lr: 0.001,
lr_variance: 0.0,
},
domain_adaptation_loss: None,
});
let effectiveness = manager.calculate_transfer_effectiveness();
assert!(effectiveness.is_some());
assert!(effectiveness.expect("operation should succeed") > 0.0); }
}