Skip to main content

sklears_neural/
transfer_learning.rs

1//! Transfer Learning Utilities
2//!
3//! This module provides comprehensive transfer learning capabilities including
4//! model freezing, layer replacement, fine-tuning strategies, and domain adaptation.
5
6use crate::layers::Layer;
7use crate::weight_init::{InitStrategy, WeightInitializer};
8use crate::NeuralResult;
9use scirs2_core::ndarray::{Array1, Array2, Array3};
10use scirs2_core::random::ChaCha8Rng;
11use scirs2_core::random::SeedableRng;
12use sklears_core::types::FloatBounds;
13use std::collections::{HashMap, HashSet};
14use std::marker::PhantomData;
15
16/// Configuration for transfer learning
17#[derive(Debug, Clone)]
18pub struct TransferConfig<T: FloatBounds> {
19    /// Layers to freeze during training
20    pub frozen_layers: HashSet<String>,
21    /// Learning rate multipliers for different layers
22    pub layer_learning_rates: HashMap<String, T>,
23    /// Fine-tuning strategy
24    pub fine_tuning_strategy: FineTuningStrategy,
25    /// Gradual unfreezing schedule
26    pub unfreeze_schedule: Option<UnfreezeSchedule>,
27    /// Domain adaptation settings
28    pub domain_adaptation: Option<DomainAdaptationConfig<T>>,
29    /// Whether to use discriminative learning rates
30    pub discriminative_lr: bool,
31    /// Base learning rate for unfrozen layers
32    pub base_learning_rate: T,
33}
34
35impl<T: FloatBounds> Default for TransferConfig<T> {
36    fn default() -> Self {
37        Self {
38            frozen_layers: HashSet::new(),
39            layer_learning_rates: HashMap::new(),
40            fine_tuning_strategy: FineTuningStrategy::FineTuneAll,
41            unfreeze_schedule: None,
42            domain_adaptation: None,
43            discriminative_lr: false,
44            base_learning_rate: T::from(1e-3).unwrap_or_else(|| T::zero()),
45        }
46    }
47}
48
49/// Fine-tuning strategies for transfer learning
50#[derive(Debug, Clone, Copy, PartialEq)]
51pub enum FineTuningStrategy {
52    /// Freeze all layers, only train new layers
53    FeatureExtraction,
54    /// Fine-tune all layers
55    FineTuneAll,
56    /// Fine-tune only top layers
57    FineTuneTop {
58        /// Number of top (output-side) layers to unfreeze for fine-tuning
59        num_layers: usize,
60    },
61    /// Gradual unfreezing strategy
62    GradualUnfreeze,
63    /// Layer-wise adaptive fine-tuning
64    LayerWiseAdaptive,
65    /// Task-specific fine-tuning
66    TaskSpecific,
67}
68
69/// Schedule for gradual unfreezing of layers
70#[derive(Debug, Clone)]
71pub struct UnfreezeSchedule {
72    /// Number of epochs between unfreezing steps
73    pub epochs_per_step: usize,
74    /// Number of layers to unfreeze at each step
75    pub layers_per_step: usize,
76    /// Start from top (last) layers or bottom (first) layers
77    pub unfreeze_direction: UnfreezeDirection,
78}
79
80/// Direction for gradual unfreezing
81#[derive(Debug, Clone, Copy, PartialEq)]
82pub enum UnfreezeDirection {
83    /// Unfreeze from top (output) layers to bottom (input) layers
84    TopToBottom,
85    /// Unfreeze from bottom (input) layers to top (output) layers
86    BottomToTop,
87    /// Unfreeze middle layers first, then expand outward
88    MiddleOut,
89}
90
91/// Domain adaptation configuration
92#[derive(Debug, Clone)]
93pub struct DomainAdaptationConfig<T: FloatBounds> {
94    /// Domain adaptation technique
95    pub technique: DomainAdaptationTechnique,
96    /// Adaptation loss weight
97    pub adaptation_weight: T,
98    /// Number of adaptation iterations
99    pub adaptation_iterations: usize,
100    /// Whether to use adversarial training
101    pub adversarial_training: bool,
102}
103
104/// Domain adaptation techniques
105#[derive(Debug, Clone, Copy, PartialEq)]
106pub enum DomainAdaptationTechnique {
107    /// Maximum Mean Discrepancy: align feature distributions via kernel mean embeddings
108    MMD,
109    /// Domain-Adversarial Neural Networks: use a gradient reversal layer to confuse a domain classifier
110    DANN,
111    /// CORrelation ALignment: minimize the covariance difference between source and target features
112    CORAL,
113    /// Adaptive Batch Normalization: re-estimate BN statistics on the target domain
114    AdaBN,
115}
116
117/// Transfer learning manager for handling model adaptation
118#[derive(Debug, Clone)]
119pub struct TransferLearningManager<T: FloatBounds> {
120    /// Transfer learning configuration
121    config: TransferConfig<T>,
122    /// Current epoch for scheduling
123    current_epoch: usize,
124    /// Layer freeze status
125    layer_freeze_status: HashMap<String, bool>,
126    /// Layer learning rate multipliers
127    layer_lr_multipliers: HashMap<String, T>,
128    /// Training history for adaptation
129    training_history: Vec<TransferMetrics<T>>,
130}
131
132impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> TransferLearningManager<T> {
133    /// Create a new transfer learning manager
134    pub fn new(config: TransferConfig<T>) -> Self {
135        let mut layer_freeze_status = HashMap::new();
136        let mut layer_lr_multipliers = HashMap::new();
137
138        // Initialize freeze status
139        for layer_name in &config.frozen_layers {
140            layer_freeze_status.insert(layer_name.clone(), true);
141        }
142
143        // Initialize learning rate multipliers
144        for (layer_name, lr_mult) in &config.layer_learning_rates {
145            layer_lr_multipliers.insert(layer_name.clone(), *lr_mult);
146        }
147
148        Self {
149            config,
150            current_epoch: 0,
151            layer_freeze_status,
152            layer_lr_multipliers,
153            training_history: Vec::new(),
154        }
155    }
156
157    /// Freeze specific layers in a model
158    pub fn freeze_layers(&mut self, layer_names: &[String]) -> NeuralResult<()> {
159        for layer_name in layer_names {
160            self.layer_freeze_status.insert(layer_name.clone(), true);
161        }
162        Ok(())
163    }
164
165    /// Unfreeze specific layers in a model
166    pub fn unfreeze_layers(&mut self, layer_names: &[String]) -> NeuralResult<()> {
167        for layer_name in layer_names {
168            self.layer_freeze_status.insert(layer_name.clone(), false);
169        }
170        Ok(())
171    }
172
173    /// Check if a layer is frozen
174    pub fn is_layer_frozen(&self, layer_name: &str) -> bool {
175        self.layer_freeze_status
176            .get(layer_name)
177            .copied()
178            .unwrap_or(false)
179    }
180
181    /// Get learning rate multiplier for a layer
182    pub fn get_layer_lr_multiplier(&self, layer_name: &str) -> T {
183        self.layer_lr_multipliers
184            .get(layer_name)
185            .copied()
186            .unwrap_or(T::one())
187    }
188
189    /// Update epoch and apply scheduling
190    pub fn update_epoch(&mut self, epoch: usize) -> NeuralResult<()> {
191        self.current_epoch = epoch;
192
193        // Apply unfreezing schedule if configured
194        if let Some(schedule) = self.config.unfreeze_schedule.clone() {
195            self.apply_unfreeze_schedule(&schedule)?;
196        }
197
198        Ok(())
199    }
200
201    /// Apply gradual unfreezing schedule
202    fn apply_unfreeze_schedule(&mut self, schedule: &UnfreezeSchedule) -> NeuralResult<()> {
203        if self.current_epoch.is_multiple_of(schedule.epochs_per_step) && self.current_epoch > 0 {
204            let step = self.current_epoch / schedule.epochs_per_step;
205            let layers_to_unfreeze = self.select_layers_for_unfreezing(schedule, step)?;
206            self.unfreeze_layers(&layers_to_unfreeze)?;
207        }
208        Ok(())
209    }
210
211    /// Select layers for unfreezing based on schedule
212    fn select_layers_for_unfreezing(
213        &self,
214        schedule: &UnfreezeSchedule,
215        step: usize,
216    ) -> NeuralResult<Vec<String>> {
217        let frozen_layers: Vec<String> = self
218            .layer_freeze_status
219            .iter()
220            .filter(|(_, &is_frozen)| is_frozen)
221            .map(|(name, _)| name.clone())
222            .collect();
223
224        let start_idx = step * schedule.layers_per_step;
225        let end_idx = ((step + 1) * schedule.layers_per_step).min(frozen_layers.len());
226
227        if start_idx >= frozen_layers.len() {
228            return Ok(Vec::new());
229        }
230
231        let selected_layers = match schedule.unfreeze_direction {
232            UnfreezeDirection::TopToBottom => frozen_layers[start_idx..end_idx].to_vec(),
233            UnfreezeDirection::BottomToTop => {
234                let mut layers = frozen_layers.clone();
235                layers.reverse();
236                layers[start_idx..end_idx].to_vec()
237            }
238            UnfreezeDirection::MiddleOut => {
239                // Unfreeze from middle outward
240                let middle = frozen_layers.len() / 2;
241                let mut selected = Vec::new();
242
243                for i in 0..schedule.layers_per_step {
244                    if step * schedule.layers_per_step + i >= frozen_layers.len() {
245                        break;
246                    }
247
248                    let offset = i / 2;
249                    if i % 2 == 0 {
250                        // Go towards end
251                        if middle + offset < frozen_layers.len() {
252                            selected.push(frozen_layers[middle + offset].clone());
253                        }
254                    } else {
255                        // Go towards beginning
256                        if offset < middle {
257                            selected.push(frozen_layers[middle - offset - 1].clone());
258                        }
259                    }
260                }
261                selected
262            }
263        };
264
265        Ok(selected_layers)
266    }
267
268    /// Apply discriminative learning rates
269    pub fn apply_discriminative_learning_rates(
270        &mut self,
271        layer_names: &[String],
272    ) -> NeuralResult<()> {
273        if !self.config.discriminative_lr {
274            return Ok(());
275        }
276
277        // Apply decreasing learning rates for lower layers
278        let num_layers = layer_names.len();
279        for (i, layer_name) in layer_names.iter().enumerate() {
280            let layer_depth = (num_layers - i - 1) as f64;
281            let lr_multiplier =
282                T::from(0.1_f64.powf(layer_depth / num_layers as f64)).unwrap_or_else(|| T::zero());
283            self.layer_lr_multipliers
284                .insert(layer_name.clone(), lr_multiplier);
285        }
286
287        Ok(())
288    }
289
290    /// Record training metrics
291    pub fn record_metrics(&mut self, metrics: TransferMetrics<T>) {
292        self.training_history.push(metrics);
293    }
294
295    /// Get training history
296    pub fn get_training_history(&self) -> &[TransferMetrics<T>] {
297        &self.training_history
298    }
299
300    /// Calculate transfer learning effectiveness
301    pub fn calculate_transfer_effectiveness(&self) -> Option<T> {
302        if self.training_history.len() < 2 {
303            return None;
304        }
305
306        let initial_loss = self.training_history[0].validation_loss;
307        let final_loss = self
308            .training_history
309            .last()
310            .expect("empty collection")
311            .validation_loss;
312
313        Some((initial_loss - final_loss) / initial_loss)
314    }
315}
316
317/// Metrics for transfer learning evaluation
318#[derive(Debug, Clone)]
319pub struct TransferMetrics<T: FloatBounds> {
320    /// Current epoch
321    pub epoch: usize,
322    /// Training loss
323    pub training_loss: T,
324    /// Validation loss
325    pub validation_loss: T,
326    /// Number of frozen layers
327    pub frozen_layer_count: usize,
328    /// Learning rate statistics
329    pub lr_stats: LearningRateStats<T>,
330    /// Domain adaptation loss (if applicable)
331    pub domain_adaptation_loss: Option<T>,
332}
333
334/// Learning rate statistics
335#[derive(Debug, Clone)]
336pub struct LearningRateStats<T: FloatBounds> {
337    /// Mean learning rate across layers
338    pub mean_lr: T,
339    /// Maximum learning rate
340    pub max_lr: T,
341    /// Minimum learning rate
342    pub min_lr: T,
343    /// Learning rate variance
344    pub lr_variance: T,
345}
346
347/// Model adapter for replacing layers during transfer learning
348pub struct ModelAdapter<T: FloatBounds> {
349    /// Layer replacement mapping
350    layer_replacements: HashMap<String, Box<dyn Layer<T>>>,
351    /// Initialization strategy for new layers
352    init_strategy: InitStrategy,
353}
354
355impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> ModelAdapter<T> {
356    /// Create a new model adapter
357    pub fn new(init_strategy: InitStrategy) -> Self {
358        Self {
359            layer_replacements: HashMap::new(),
360            init_strategy,
361        }
362    }
363
364    /// Add a layer replacement
365    pub fn replace_layer(&mut self, layer_name: String, new_layer: Box<dyn Layer<T>>) {
366        self.layer_replacements.insert(layer_name, new_layer);
367    }
368
369    /// Remove the final classification layer and add a new one
370    pub fn replace_classifier(
371        &mut self,
372        num_classes: usize,
373        hidden_size: usize,
374        _layer_name: String,
375    ) -> NeuralResult<()> {
376        let mut rng = ChaCha8Rng::seed_from_u64(42);
377        let initializer: WeightInitializer<T> = WeightInitializer::new(self.init_strategy);
378
379        // Create new classification layer (simple dense layer implementation)
380        let _weights = initializer.initialize_2d(&mut rng, (hidden_size, num_classes))?;
381        let _bias: Array1<T> = Array1::zeros(num_classes);
382
383        // For now, we'll create a placeholder - in practice this would be a proper Dense layer
384        // self.layer_replacements.insert(layer_name, Box::new(DenseLayer::new(weights, bias)));
385
386        Ok(())
387    }
388
389    /// Apply layer replacements to a model
390    pub fn apply_replacements<M>(&self, model: &mut M) -> NeuralResult<()>
391    where
392        M: HasReplaceableLayer<T>,
393    {
394        for (layer_name, replacement_layer) in &self.layer_replacements {
395            model.replace_layer(layer_name, replacement_layer.as_ref())?;
396        }
397        Ok(())
398    }
399}
400
401impl<T: FloatBounds> std::fmt::Debug for ModelAdapter<T> {
402    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
403        f.debug_struct("ModelAdapter")
404            .field(
405                "layer_replacements",
406                &format!("{} layers", self.layer_replacements.len()),
407            )
408            .field("init_strategy", &self.init_strategy)
409            .finish()
410    }
411}
412
413/// Trait for models that support layer replacement
414pub trait HasReplaceableLayer<T: FloatBounds> {
415    /// Replace a layer in the model
416    fn replace_layer(&mut self, layer_name: &str, new_layer: &dyn Layer<T>) -> NeuralResult<()>;
417
418    /// Get layer names
419    fn get_layer_names(&self) -> Vec<String>;
420}
421
422/// Feature extractor for using pre-trained models as feature extractors
423#[derive(Debug, Clone)]
424pub struct FeatureExtractor<T: FloatBounds> {
425    /// Extract features up to this layer
426    extract_layer: String,
427    /// Whether to apply global pooling
428    global_pooling: bool,
429    /// Pooling strategy
430    pooling_strategy: PoolingStrategy,
431    /// Phantom data for type parameter
432    _phantom: PhantomData<T>,
433}
434
435impl<T: FloatBounds> FeatureExtractor<T> {
436    /// Create a new feature extractor
437    pub fn new(
438        extract_layer: String,
439        global_pooling: bool,
440        pooling_strategy: PoolingStrategy,
441    ) -> Self {
442        Self {
443            extract_layer,
444            global_pooling,
445            pooling_strategy,
446            _phantom: PhantomData,
447        }
448    }
449
450    /// Extract features from input using the specified layer
451    pub fn extract_features<M>(&self, model: &M, input: &Array3<T>) -> NeuralResult<Array2<T>>
452    where
453        M: HasFeatureExtraction<T>,
454    {
455        let features = model.extract_features_at_layer(&self.extract_layer, input)?;
456
457        if self.global_pooling {
458            self.apply_global_pooling(&features)
459        } else {
460            // Flatten features
461            let (batch_size, _, _) = features.dim();
462            let flattened_size = features.len() / batch_size;
463            Ok(features
464                .into_shape_with_order((batch_size, flattened_size))
465                .expect("array shape error"))
466        }
467    }
468
469    /// Apply global pooling to features
470    fn apply_global_pooling(&self, features: &Array3<T>) -> NeuralResult<Array2<T>> {
471        let (batch_size, _, num_features) = features.dim();
472        let mut pooled = Array2::zeros((batch_size, num_features));
473
474        for b in 0..batch_size {
475            for f in 0..num_features {
476                let feature_map = features.slice(scirs2_core::ndarray::s![b, .., f]);
477                let pooled_value = match self.pooling_strategy {
478                    PoolingStrategy::Mean => feature_map.mean().unwrap_or(T::zero()),
479                    PoolingStrategy::Max => feature_map.fold(T::zero(), |acc, &x| acc.max(x)),
480                    PoolingStrategy::Min => feature_map.fold(T::zero(), |acc, &x| acc.min(x)),
481                };
482                pooled[[b, f]] = pooled_value;
483            }
484        }
485
486        Ok(pooled)
487    }
488}
489
490/// Pooling strategies for feature extraction
491#[derive(Debug, Clone, Copy, PartialEq)]
492pub enum PoolingStrategy {
493    /// Mean pooling
494    Mean,
495    /// Max pooling
496    Max,
497    /// Min pooling
498    Min,
499}
500
501/// Trait for models that support feature extraction
502pub trait HasFeatureExtraction<T: FloatBounds> {
503    /// Extract features at a specific layer
504    fn extract_features_at_layer(
505        &self,
506        layer_name: &str,
507        input: &Array3<T>,
508    ) -> NeuralResult<Array3<T>>;
509}
510
511/// Domain adaptation utilities
512#[derive(Debug, Clone)]
513pub struct DomainAdapter<T: FloatBounds> {
514    /// Domain adaptation configuration
515    config: DomainAdaptationConfig<T>,
516    /// Source domain statistics
517    source_stats: Option<DomainStatistics<T>>,
518    /// Target domain statistics
519    target_stats: Option<DomainStatistics<T>>,
520}
521
522impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> DomainAdapter<T> {
523    /// Create a new domain adapter
524    pub fn new(config: DomainAdaptationConfig<T>) -> Self {
525        Self {
526            config,
527            source_stats: None,
528            target_stats: None,
529        }
530    }
531
532    /// Compute domain statistics
533    pub fn compute_domain_statistics(
534        &mut self,
535        source_data: &Array3<T>,
536        target_data: &Array3<T>,
537    ) -> NeuralResult<()> {
538        self.source_stats = Some(self.compute_statistics(source_data)?);
539        self.target_stats = Some(self.compute_statistics(target_data)?);
540        Ok(())
541    }
542
543    /// Apply domain adaptation
544    pub fn adapt_features(&self, features: &Array3<T>, is_source: bool) -> NeuralResult<Array3<T>> {
545        match self.config.technique {
546            DomainAdaptationTechnique::CORAL => self.apply_coral_adaptation(features, is_source),
547            DomainAdaptationTechnique::AdaBN => self.apply_adaptive_batch_norm(features, is_source),
548            _ => {
549                // Other techniques require more complex implementation
550                Ok(features.clone())
551            }
552        }
553    }
554
555    /// Compute statistics for a domain
556    fn compute_statistics(&self, data: &Array3<T>) -> NeuralResult<DomainStatistics<T>> {
557        let (batch_size, seq_len, feature_dim) = data.dim();
558        let total_samples = batch_size * seq_len;
559
560        // Compute mean
561        let mut mean = Array1::zeros(feature_dim);
562        for b in 0..batch_size {
563            for s in 0..seq_len {
564                for f in 0..feature_dim {
565                    mean[f] += data[[b, s, f]];
566                }
567            }
568        }
569        mean /= T::from(total_samples).unwrap_or_else(|| T::zero());
570
571        // Compute covariance
572        let mut covariance = Array2::zeros((feature_dim, feature_dim));
573        for b in 0..batch_size {
574            for s in 0..seq_len {
575                for i in 0..feature_dim {
576                    for j in 0..feature_dim {
577                        let diff_i = data[[b, s, i]] - mean[i];
578                        let diff_j = data[[b, s, j]] - mean[j];
579                        covariance[[i, j]] += diff_i * diff_j;
580                    }
581                }
582            }
583        }
584        covariance /= T::from(total_samples - 1).unwrap_or_else(|| T::zero());
585
586        Ok(DomainStatistics { mean, covariance })
587    }
588
589    /// Apply CORAL (Correlation Alignment) adaptation
590    fn apply_coral_adaptation(
591        &self,
592        features: &Array3<T>,
593        is_source: bool,
594    ) -> NeuralResult<Array3<T>> {
595        let stats = if is_source {
596            &self.source_stats
597        } else {
598            &self.target_stats
599        };
600
601        if let Some(domain_stats) = stats {
602            // For CORAL, we would typically align second-order statistics
603            // This is a simplified implementation
604            let (batch_size, seq_len, feature_dim) = features.dim();
605            let mut adapted = features.clone();
606
607            // Center the features
608            for b in 0..batch_size {
609                for s in 0..seq_len {
610                    for f in 0..feature_dim {
611                        adapted[[b, s, f]] -= domain_stats.mean[f];
612                    }
613                }
614            }
615
616            Ok(adapted)
617        } else {
618            Ok(features.clone())
619        }
620    }
621
622    /// Apply adaptive batch normalization
623    fn apply_adaptive_batch_norm(
624        &self,
625        features: &Array3<T>,
626        _is_source: bool,
627    ) -> NeuralResult<Array3<T>> {
628        // Simplified adaptive batch normalization
629        let (batch_size, seq_len, feature_dim) = features.dim();
630        let mut normalized = Array3::zeros((batch_size, seq_len, feature_dim));
631
632        for f in 0..feature_dim {
633            // Compute statistics for this feature across all samples
634            let mut sum = T::zero();
635            let mut count = 0;
636
637            for b in 0..batch_size {
638                for s in 0..seq_len {
639                    sum += features[[b, s, f]];
640                    count += 1;
641                }
642            }
643
644            let mean = sum / T::from(count).unwrap_or_else(|| T::zero());
645
646            let mut variance_sum = T::zero();
647            for b in 0..batch_size {
648                for s in 0..seq_len {
649                    let diff = features[[b, s, f]] - mean;
650                    variance_sum += diff * diff;
651                }
652            }
653
654            let variance = variance_sum / T::from(count).unwrap_or_else(|| T::zero());
655            let std = (variance + T::from(1e-5).unwrap_or_else(|| T::zero())).sqrt();
656
657            // Normalize
658            for b in 0..batch_size {
659                for s in 0..seq_len {
660                    normalized[[b, s, f]] = (features[[b, s, f]] - mean) / std;
661                }
662            }
663        }
664
665        Ok(normalized)
666    }
667}
668
669/// Domain statistics for adaptation
670#[derive(Debug, Clone)]
671pub struct DomainStatistics<T: FloatBounds> {
672    /// Feature means
673    pub mean: Array1<T>,
674    /// Feature covariance matrix
675    pub covariance: Array2<T>,
676}
677
678#[allow(non_snake_case)]
679#[cfg(test)]
680mod tests {
681    use super::*;
682    use scirs2_core::essentials::Normal;
683    use scirs2_core::ndarray::Array3;
684    use scirs2_core::random::thread_rng;
685
686    #[test]
687    fn test_transfer_learning_manager_creation() {
688        let config = TransferConfig::default();
689        let manager = TransferLearningManager::<f64>::new(config);
690        assert_eq!(manager.current_epoch, 0);
691    }
692
693    #[test]
694    fn test_layer_freezing() {
695        let mut config = TransferConfig::<f64>::default();
696        config.frozen_layers.insert("layer1".to_string());
697
698        let mut manager = TransferLearningManager::new(config);
699        assert!(manager.is_layer_frozen("layer1"));
700        assert!(!manager.is_layer_frozen("layer2"));
701
702        manager
703            .unfreeze_layers(&["layer1".to_string()])
704            .expect("operation should succeed");
705        assert!(!manager.is_layer_frozen("layer1"));
706    }
707
708    #[test]
709    fn test_learning_rate_multipliers() {
710        let mut config = TransferConfig::<f64>::default();
711        config
712            .layer_learning_rates
713            .insert("layer1".to_string(), 0.5);
714
715        let manager = TransferLearningManager::new(config);
716        assert_eq!(manager.get_layer_lr_multiplier("layer1"), 0.5);
717        assert_eq!(manager.get_layer_lr_multiplier("layer2"), 1.0);
718    }
719
720    #[test]
721    fn test_model_adapter_creation() {
722        let adapter = ModelAdapter::<f64>::new(InitStrategy::XavierUniform);
723        assert_eq!(adapter.layer_replacements.len(), 0);
724    }
725
726    #[test]
727    fn test_feature_extractor() {
728        let extractor =
729            FeatureExtractor::<f64>::new("conv_layer".to_string(), true, PoolingStrategy::Mean);
730        assert_eq!(extractor.extract_layer, "conv_layer");
731        assert!(extractor.global_pooling);
732    }
733
734    #[test]
735    fn test_domain_adapter() {
736        let config = DomainAdaptationConfig {
737            technique: DomainAdaptationTechnique::CORAL,
738            adaptation_weight: 0.1,
739            adaptation_iterations: 100,
740            adversarial_training: false,
741        };
742
743        let adapter = DomainAdapter::<f64>::new(config);
744        assert!(adapter.source_stats.is_none());
745        assert!(adapter.target_stats.is_none());
746    }
747
748    #[test]
749    fn test_domain_statistics_computation() {
750        let config = DomainAdaptationConfig {
751            technique: DomainAdaptationTechnique::CORAL,
752            adaptation_weight: 0.1,
753            adaptation_iterations: 100,
754            adversarial_training: false,
755        };
756
757        let mut adapter = DomainAdapter::new(config);
758
759        let source_data = Array3::from_shape_fn((10, 5, 8), |_| {
760            let mut rng = thread_rng();
761            rng.sample(Normal::new(0.0, 1.0).expect("construction should succeed"))
762        });
763        let target_data = Array3::from_shape_fn((10, 5, 8), |_| {
764            let mut rng = thread_rng();
765            rng.sample(Normal::new(0.0, 1.0).expect("construction should succeed"))
766        });
767
768        let result = adapter.compute_domain_statistics(&source_data, &target_data);
769        assert!(result.is_ok());
770        assert!(adapter.source_stats.is_some());
771        assert!(adapter.target_stats.is_some());
772    }
773
774    #[test]
775    fn test_unfreeze_schedule() {
776        let schedule = UnfreezeSchedule {
777            epochs_per_step: 5,
778            layers_per_step: 2,
779            unfreeze_direction: UnfreezeDirection::TopToBottom,
780        };
781
782        let mut config = TransferConfig::<f64> {
783            unfreeze_schedule: Some(schedule),
784            ..Default::default()
785        };
786
787        // Add some frozen layers
788        for i in 0..6 {
789            config.frozen_layers.insert(format!("layer_{}", i));
790        }
791
792        let mut manager = TransferLearningManager::new(config);
793
794        // Update to epoch 5 (should trigger unfreezing)
795        manager.update_epoch(5).expect("operation should succeed");
796
797        // Some layers should have been unfrozen
798        let frozen_count = manager.layer_freeze_status.values().filter(|&&v| v).count();
799        assert!(frozen_count < 6);
800    }
801
802    #[test]
803    fn test_transfer_effectiveness_calculation() {
804        let config = TransferConfig::default();
805        let mut manager = TransferLearningManager::new(config);
806
807        // Add some training history
808        manager.record_metrics(TransferMetrics {
809            epoch: 0,
810            training_loss: 1.0,
811            validation_loss: 1.0,
812            frozen_layer_count: 5,
813            lr_stats: LearningRateStats {
814                mean_lr: 0.001,
815                max_lr: 0.001,
816                min_lr: 0.001,
817                lr_variance: 0.0,
818            },
819            domain_adaptation_loss: None,
820        });
821
822        manager.record_metrics(TransferMetrics {
823            epoch: 10,
824            training_loss: 0.5,
825            validation_loss: 0.6,
826            frozen_layer_count: 3,
827            lr_stats: LearningRateStats {
828                mean_lr: 0.001,
829                max_lr: 0.001,
830                min_lr: 0.001,
831                lr_variance: 0.0,
832            },
833            domain_adaptation_loss: None,
834        });
835
836        let effectiveness = manager.calculate_transfer_effectiveness();
837        assert!(effectiveness.is_some());
838        assert!(effectiveness.expect("operation should succeed") > 0.0); // Should show improvement
839    }
840}