1use super::config::ArchitectureSpec;
7use super::config::*;
8use super::features::DistributionStats;
9use crate::applications::ApplicationResult;
10use std::collections::HashMap;
11use std::time::Instant;
12
13pub struct TransferLearner {
15 pub source_domains: Vec<SourceDomain>,
17 pub similarity_analyzer: DomainSimilarityAnalyzer,
19 pub transfer_strategies: Vec<TransferStrategy>,
21 pub adaptation_mechanisms: Vec<AdaptationMechanism>,
23}
24
25#[derive(Debug)]
27pub struct SourceDomain {
28 pub id: String,
30 pub characteristics: DomainCharacteristics,
32 pub models: Vec<TransferableModel>,
34 pub transfer_history: Vec<TransferRecord>,
36}
37
38#[derive(Debug, Clone)]
40pub struct DomainCharacteristics {
41 pub feature_distribution: DistributionStats,
43 pub label_distribution: DistributionStats,
45 pub task_complexity: f64,
47 pub data_size: usize,
49 pub noise_level: f64,
51}
52
53#[derive(Debug)]
55pub struct TransferableModel {
56 pub id: String,
58 pub architecture: ArchitectureSpec,
60 pub weights: Vec<f64>,
62 pub source_performance: f64,
64 pub transferability_score: f64,
66}
67
68#[derive(Debug, Clone)]
70pub struct TransferRecord {
71 pub timestamp: Instant,
73 pub target_domain: String,
75 pub strategy: TransferStrategy,
77 pub performance_improvement: f64,
79 pub success: bool,
81}
82
83#[derive(Debug)]
85pub struct DomainSimilarityAnalyzer {
86 pub metrics: Vec<SimilarityMetric>,
88 pub similarity_cache: HashMap<(String, String), f64>,
90 pub methods: Vec<SimilarityMethod>,
92}
93
94#[derive(Debug, Clone, PartialEq, Eq)]
96pub enum SimilarityMetric {
97 FeatureSimilarity,
99 TaskSimilarity,
101 DataDistributionSimilarity,
103 PerformanceCorrelation,
105 StructuralSimilarity,
107}
108
109#[derive(Debug, Clone, PartialEq, Eq)]
111pub enum SimilarityMethod {
112 Cosine,
114 Euclidean,
116 Wasserstein,
118 MaximumMeanDiscrepancy,
120 Kernel(String),
122}
123
124#[derive(Debug, Clone, PartialEq, Eq)]
126pub enum TransferStrategy {
127 FeatureTransfer,
129 ParameterTransfer,
131 InstanceTransfer,
133 RelationalTransfer,
135 MultiTaskLearning,
137 DomainAdaptation,
139}
140
141#[derive(Debug, Clone, PartialEq, Eq)]
143pub enum AdaptationMechanism {
144 FineTuning,
146 DomainAdversarial,
148 GradualUnfreezing,
150 KnowledgeDistillation,
152 ProgressiveTraining,
154}
155
156impl TransferLearner {
157 #[must_use]
158 pub fn new() -> Self {
159 Self {
160 source_domains: Vec::new(),
161 similarity_analyzer: DomainSimilarityAnalyzer {
162 metrics: vec![SimilarityMetric::FeatureSimilarity],
163 similarity_cache: HashMap::new(),
164 methods: vec![SimilarityMethod::Cosine],
165 },
166 transfer_strategies: vec![TransferStrategy::ParameterTransfer],
167 adaptation_mechanisms: vec![AdaptationMechanism::FineTuning],
168 }
169 }
170
171 pub fn add_source_domain(&mut self, domain: SourceDomain) {
173 self.source_domains.push(domain);
174 }
175
176 #[must_use]
178 pub fn find_similar_domain(
179 &self,
180 target_characteristics: &DomainCharacteristics,
181 ) -> Option<&SourceDomain> {
182 let mut best_domain = None;
183 let mut best_similarity = 0.0;
184
185 for domain in &self.source_domains {
186 let similarity =
187 self.calculate_domain_similarity(&domain.characteristics, target_characteristics);
188 if similarity > best_similarity {
189 best_similarity = similarity;
190 best_domain = Some(domain);
191 }
192 }
193
194 best_domain
195 }
196
197 fn calculate_domain_similarity(
199 &self,
200 source: &DomainCharacteristics,
201 target: &DomainCharacteristics,
202 ) -> f64 {
203 let complexity_sim = 1.0 - (source.task_complexity - target.task_complexity).abs();
205 let size_sim =
206 1.0 - ((source.data_size as f64).ln() - (target.data_size as f64).ln()).abs() / 10.0;
207 let noise_sim = 1.0 - (source.noise_level - target.noise_level).abs();
208
209 (complexity_sim * 0.4 + size_sim * 0.3 + noise_sim * 0.3)
211 .max(0.0)
212 .min(1.0)
213 }
214
215 pub fn transfer_knowledge(
217 &mut self,
218 source_domain_id: &str,
219 target_domain: &str,
220 strategy: TransferStrategy,
221 ) -> ApplicationResult<TransferResult> {
222 let source_domain = self
224 .source_domains
225 .iter()
226 .find(|d| d.id == source_domain_id)
227 .ok_or_else(|| {
228 crate::applications::ApplicationError::InvalidConfiguration(format!(
229 "Source domain {source_domain_id} not found"
230 ))
231 })?;
232
233 let performance_improvement = match strategy {
235 TransferStrategy::ParameterTransfer => 0.15,
236 TransferStrategy::FeatureTransfer => 0.12,
237 TransferStrategy::DomainAdaptation => 0.18,
238 _ => 0.10,
239 };
240
241 let record = TransferRecord {
243 timestamp: Instant::now(),
244 target_domain: target_domain.to_string(),
245 strategy: strategy.clone(),
246 performance_improvement,
247 success: performance_improvement > 0.05,
248 };
249
250 Ok(TransferResult {
253 success: record.success,
254 performance_improvement: record.performance_improvement,
255 transfer_method: strategy,
256 confidence: 0.8,
257 })
258 }
259
260 #[must_use]
262 pub fn get_transfer_statistics(&self) -> TransferStatistics {
263 let mut total_transfers = 0;
264 let mut successful_transfers = 0;
265 let mut total_improvement = 0.0;
266
267 for domain in &self.source_domains {
268 for record in &domain.transfer_history {
269 total_transfers += 1;
270 if record.success {
271 successful_transfers += 1;
272 total_improvement += record.performance_improvement;
273 }
274 }
275 }
276
277 let success_rate = if total_transfers > 0 {
278 successful_transfers as f64 / total_transfers as f64
279 } else {
280 0.0
281 };
282
283 let avg_improvement = if successful_transfers > 0 {
284 total_improvement / successful_transfers as f64
285 } else {
286 0.0
287 };
288
289 TransferStatistics {
290 total_transfers,
291 successful_transfers,
292 success_rate,
293 average_improvement: avg_improvement,
294 }
295 }
296}
297
298#[derive(Debug, Clone)]
300pub struct TransferResult {
301 pub success: bool,
303 pub performance_improvement: f64,
305 pub transfer_method: TransferStrategy,
307 pub confidence: f64,
309}
310
311#[derive(Debug, Clone)]
313pub struct TransferStatistics {
314 pub total_transfers: usize,
316 pub successful_transfers: usize,
318 pub success_rate: f64,
320 pub average_improvement: f64,
322}
323
324#[cfg(test)]
325mod tests {
326 use super::*;
327 use crate::meta_learning::config::{
328 ActivationFunction, ConnectionPattern, LayerSpec, LayerType, OptimizationSettings,
329 OptimizerType, RegularizationConfig,
330 };
331 use std::time::Duration;
332
333 #[test]
334 fn test_transfer_learner_creation() {
335 let learner = TransferLearner::new();
336 assert_eq!(learner.source_domains.len(), 0);
337 assert_eq!(learner.transfer_strategies.len(), 1);
338 assert_eq!(learner.adaptation_mechanisms.len(), 1);
339 }
340
341 #[test]
342 fn test_domain_similarity() {
343 let learner = TransferLearner::new();
344
345 let source = DomainCharacteristics {
346 feature_distribution: DistributionStats::default(),
347 label_distribution: DistributionStats::default(),
348 task_complexity: 0.5,
349 data_size: 1000,
350 noise_level: 0.1,
351 };
352
353 let target = DomainCharacteristics {
354 feature_distribution: DistributionStats::default(),
355 label_distribution: DistributionStats::default(),
356 task_complexity: 0.6,
357 data_size: 1200,
358 noise_level: 0.15,
359 };
360
361 let similarity = learner.calculate_domain_similarity(&source, &target);
362 assert!(similarity > 0.0);
363 assert!(similarity <= 1.0);
364 }
365
366 #[test]
367 fn test_source_domain_addition() {
368 let mut learner = TransferLearner::new();
369
370 let domain = SourceDomain {
371 id: "test_domain".to_string(),
372 characteristics: DomainCharacteristics {
373 feature_distribution: DistributionStats::default(),
374 label_distribution: DistributionStats::default(),
375 task_complexity: 0.5,
376 data_size: 1000,
377 noise_level: 0.1,
378 },
379 models: vec![TransferableModel {
380 id: "test_model".to_string(),
381 architecture: ArchitectureSpec {
382 layers: vec![LayerSpec {
383 layer_type: LayerType::Dense,
384 input_dim: 10,
385 output_dim: 5,
386 activation: ActivationFunction::ReLU,
387 dropout: 0.1,
388 parameters: HashMap::new(),
389 }],
390 connections: ConnectionPattern::Sequential,
391 optimization: OptimizationSettings {
392 optimizer: OptimizerType::Adam,
393 learning_rate: 0.001,
394 batch_size: 32,
395 epochs: 100,
396 regularization: RegularizationConfig {
397 l1_weight: 0.0,
398 l2_weight: 0.01,
399 dropout: 0.1,
400 batch_norm: true,
401 early_stopping: true,
402 },
403 },
404 },
405 weights: vec![0.1, 0.2, 0.3],
406 source_performance: 0.9,
407 transferability_score: 0.8,
408 }],
409 transfer_history: Vec::new(),
410 };
411
412 learner.add_source_domain(domain);
413 assert_eq!(learner.source_domains.len(), 1);
414 }
415
416 #[test]
417 fn test_transfer_knowledge() {
418 let mut learner = TransferLearner::new();
419
420 let domain = SourceDomain {
422 id: "source_domain".to_string(),
423 characteristics: DomainCharacteristics {
424 feature_distribution: DistributionStats::default(),
425 label_distribution: DistributionStats::default(),
426 task_complexity: 0.5,
427 data_size: 1000,
428 noise_level: 0.1,
429 },
430 models: Vec::new(),
431 transfer_history: Vec::new(),
432 };
433
434 learner.add_source_domain(domain);
435
436 let result = learner.transfer_knowledge(
438 "source_domain",
439 "target_domain",
440 TransferStrategy::ParameterTransfer,
441 );
442
443 assert!(result.is_ok());
444 let transfer_result = result.expect("Transfer knowledge should succeed");
445 assert!(transfer_result.performance_improvement > 0.0);
446 }
447
448 #[test]
449 fn test_transfer_statistics() {
450 let learner = TransferLearner::new();
451 let stats = learner.get_transfer_statistics();
452
453 assert_eq!(stats.total_transfers, 0);
454 assert_eq!(stats.successful_transfers, 0);
455 assert_eq!(stats.success_rate, 0.0);
456 assert_eq!(stats.average_improvement, 0.0);
457 }
458
459 #[test]
460 fn test_similarity_metrics() {
461 assert_eq!(
462 SimilarityMetric::FeatureSimilarity,
463 SimilarityMetric::FeatureSimilarity
464 );
465 assert_ne!(
466 SimilarityMetric::FeatureSimilarity,
467 SimilarityMetric::TaskSimilarity
468 );
469 }
470
471 #[test]
472 fn test_transfer_strategies() {
473 assert_eq!(
474 TransferStrategy::ParameterTransfer,
475 TransferStrategy::ParameterTransfer
476 );
477 assert_ne!(
478 TransferStrategy::ParameterTransfer,
479 TransferStrategy::FeatureTransfer
480 );
481 }
482
483 #[test]
484 fn test_adaptation_mechanisms() {
485 assert_eq!(
486 AdaptationMechanism::FineTuning,
487 AdaptationMechanism::FineTuning
488 );
489 assert_ne!(
490 AdaptationMechanism::FineTuning,
491 AdaptationMechanism::DomainAdversarial
492 );
493 }
494}