1use 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#[derive(Debug, Clone)]
18pub struct TransferConfig<T: FloatBounds> {
19 pub frozen_layers: HashSet<String>,
21 pub layer_learning_rates: HashMap<String, T>,
23 pub fine_tuning_strategy: FineTuningStrategy,
25 pub unfreeze_schedule: Option<UnfreezeSchedule>,
27 pub domain_adaptation: Option<DomainAdaptationConfig<T>>,
29 pub discriminative_lr: bool,
31 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#[derive(Debug, Clone, Copy, PartialEq)]
51pub enum FineTuningStrategy {
52 FeatureExtraction,
54 FineTuneAll,
56 FineTuneTop {
58 num_layers: usize,
60 },
61 GradualUnfreeze,
63 LayerWiseAdaptive,
65 TaskSpecific,
67}
68
69#[derive(Debug, Clone)]
71pub struct UnfreezeSchedule {
72 pub epochs_per_step: usize,
74 pub layers_per_step: usize,
76 pub unfreeze_direction: UnfreezeDirection,
78}
79
80#[derive(Debug, Clone, Copy, PartialEq)]
82pub enum UnfreezeDirection {
83 TopToBottom,
85 BottomToTop,
87 MiddleOut,
89}
90
91#[derive(Debug, Clone)]
93pub struct DomainAdaptationConfig<T: FloatBounds> {
94 pub technique: DomainAdaptationTechnique,
96 pub adaptation_weight: T,
98 pub adaptation_iterations: usize,
100 pub adversarial_training: bool,
102}
103
104#[derive(Debug, Clone, Copy, PartialEq)]
106pub enum DomainAdaptationTechnique {
107 MMD,
109 DANN,
111 CORAL,
113 AdaBN,
115}
116
117#[derive(Debug, Clone)]
119pub struct TransferLearningManager<T: FloatBounds> {
120 config: TransferConfig<T>,
122 current_epoch: usize,
124 layer_freeze_status: HashMap<String, bool>,
126 layer_lr_multipliers: HashMap<String, T>,
128 training_history: Vec<TransferMetrics<T>>,
130}
131
132impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> TransferLearningManager<T> {
133 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 for layer_name in &config.frozen_layers {
140 layer_freeze_status.insert(layer_name.clone(), true);
141 }
142
143 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 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 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 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 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 pub fn update_epoch(&mut self, epoch: usize) -> NeuralResult<()> {
191 self.current_epoch = epoch;
192
193 if let Some(schedule) = self.config.unfreeze_schedule.clone() {
195 self.apply_unfreeze_schedule(&schedule)?;
196 }
197
198 Ok(())
199 }
200
201 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 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 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 if middle + offset < frozen_layers.len() {
252 selected.push(frozen_layers[middle + offset].clone());
253 }
254 } else {
255 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 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 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 pub fn record_metrics(&mut self, metrics: TransferMetrics<T>) {
292 self.training_history.push(metrics);
293 }
294
295 pub fn get_training_history(&self) -> &[TransferMetrics<T>] {
297 &self.training_history
298 }
299
300 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#[derive(Debug, Clone)]
319pub struct TransferMetrics<T: FloatBounds> {
320 pub epoch: usize,
322 pub training_loss: T,
324 pub validation_loss: T,
326 pub frozen_layer_count: usize,
328 pub lr_stats: LearningRateStats<T>,
330 pub domain_adaptation_loss: Option<T>,
332}
333
334#[derive(Debug, Clone)]
336pub struct LearningRateStats<T: FloatBounds> {
337 pub mean_lr: T,
339 pub max_lr: T,
341 pub min_lr: T,
343 pub lr_variance: T,
345}
346
347pub struct ModelAdapter<T: FloatBounds> {
349 layer_replacements: HashMap<String, Box<dyn Layer<T>>>,
351 init_strategy: InitStrategy,
353}
354
355impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> ModelAdapter<T> {
356 pub fn new(init_strategy: InitStrategy) -> Self {
358 Self {
359 layer_replacements: HashMap::new(),
360 init_strategy,
361 }
362 }
363
364 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 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 let _weights = initializer.initialize_2d(&mut rng, (hidden_size, num_classes))?;
381 let _bias: Array1<T> = Array1::zeros(num_classes);
382
383 Ok(())
387 }
388
389 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
413pub trait HasReplaceableLayer<T: FloatBounds> {
415 fn replace_layer(&mut self, layer_name: &str, new_layer: &dyn Layer<T>) -> NeuralResult<()>;
417
418 fn get_layer_names(&self) -> Vec<String>;
420}
421
422#[derive(Debug, Clone)]
424pub struct FeatureExtractor<T: FloatBounds> {
425 extract_layer: String,
427 global_pooling: bool,
429 pooling_strategy: PoolingStrategy,
431 _phantom: PhantomData<T>,
433}
434
435impl<T: FloatBounds> FeatureExtractor<T> {
436 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 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 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 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#[derive(Debug, Clone, Copy, PartialEq)]
492pub enum PoolingStrategy {
493 Mean,
495 Max,
497 Min,
499}
500
501pub trait HasFeatureExtraction<T: FloatBounds> {
503 fn extract_features_at_layer(
505 &self,
506 layer_name: &str,
507 input: &Array3<T>,
508 ) -> NeuralResult<Array3<T>>;
509}
510
511#[derive(Debug, Clone)]
513pub struct DomainAdapter<T: FloatBounds> {
514 config: DomainAdaptationConfig<T>,
516 source_stats: Option<DomainStatistics<T>>,
518 target_stats: Option<DomainStatistics<T>>,
520}
521
522impl<T: FloatBounds + scirs2_core::ndarray::ScalarOperand> DomainAdapter<T> {
523 pub fn new(config: DomainAdaptationConfig<T>) -> Self {
525 Self {
526 config,
527 source_stats: None,
528 target_stats: None,
529 }
530 }
531
532 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 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 Ok(features.clone())
551 }
552 }
553 }
554
555 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 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 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 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 let (batch_size, seq_len, feature_dim) = features.dim();
605 let mut adapted = features.clone();
606
607 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 fn apply_adaptive_batch_norm(
624 &self,
625 features: &Array3<T>,
626 _is_source: bool,
627 ) -> NeuralResult<Array3<T>> {
628 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 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 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#[derive(Debug, Clone)]
671pub struct DomainStatistics<T: FloatBounds> {
672 pub mean: Array1<T>,
674 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 for i in 0..6 {
789 config.frozen_layers.insert(format!("layer_{}", i));
790 }
791
792 let mut manager = TransferLearningManager::new(config);
793
794 manager.update_epoch(5).expect("operation should succeed");
796
797 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 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); }
840}