1use crate::{EmbeddingModel, Vector};
7use anyhow::{anyhow, Result};
8use scirs2_core::random::{Random, RngExt};
9use serde::{Deserialize, Serialize};
10use std::collections::{HashMap, HashSet};
11
12pub struct CrossDomainTransferManager {
14 source_domains: HashMap<String, DomainModel>,
16 target_domains: HashMap<String, DomainSpecification>,
18 transfer_strategies: Vec<TransferStrategy>,
20 transfer_metrics: Vec<TransferMetric>,
22 config: TransferConfig,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
28pub struct TransferConfig {
29 pub enable_domain_adaptation: bool,
31 pub use_adversarial_alignment: bool,
33 pub max_alignment_iterations: usize,
35 pub adaptation_learning_rate: f64,
37 pub min_domain_similarity: f64,
39 pub enable_entity_linking: bool,
41 pub evaluation_sample_size: usize,
43}
44
45impl Default for TransferConfig {
46 fn default() -> Self {
47 Self {
48 enable_domain_adaptation: true,
49 use_adversarial_alignment: true,
50 max_alignment_iterations: 100,
51 adaptation_learning_rate: 0.001,
52 min_domain_similarity: 0.3,
53 enable_entity_linking: true,
54 evaluation_sample_size: 1000,
55 }
56 }
57}
58
59pub struct DomainModel {
61 pub domain_id: String,
63 pub model: Box<dyn EmbeddingModel + Send + Sync>,
65 pub characteristics: DomainCharacteristics,
67 pub entity_mappings: HashMap<String, String>,
69 pub vocabulary: HashSet<String>,
71}
72
73#[derive(Debug, Clone, Serialize, Deserialize)]
75pub struct DomainCharacteristics {
76 pub domain_type: String,
78 pub language: String,
80 pub entity_types: Vec<String>,
82 pub relation_types: Vec<String>,
84 pub size_metrics: DomainSizeMetrics,
86 pub complexity_metrics: DomainComplexityMetrics,
88}
89
90#[derive(Debug, Clone, Serialize, Deserialize)]
92pub struct DomainSizeMetrics {
93 pub num_entities: usize,
95 pub num_relations: usize,
97 pub num_triples: usize,
99 pub avg_entity_degree: f64,
101 pub graph_density: f64,
103}
104
105#[derive(Debug, Clone, Serialize, Deserialize)]
107pub struct DomainComplexityMetrics {
108 pub entity_type_diversity: usize,
110 pub relation_type_diversity: usize,
112 pub hierarchical_depth: usize,
114 pub semantic_diversity: f64,
116 pub structural_complexity: f64,
118}
119
120#[derive(Debug, Clone, Serialize, Deserialize)]
122pub struct DomainSpecification {
123 pub domain_id: String,
125 pub characteristics: DomainCharacteristics,
127 pub training_data: Vec<(String, String, String)>,
129 pub validation_data: Vec<(String, String, String)>,
131 pub test_data: Vec<(String, String, String)>,
133}
134
135#[derive(Debug, Clone, Serialize, Deserialize)]
137pub enum TransferStrategy {
138 DirectTransfer,
140 FineTuning {
142 learning_rate: f64,
143 epochs: usize,
144 freeze_layers: Vec<String>,
145 },
146 DomainAdaptation {
148 alignment_method: AlignmentMethod,
149 regularization_strength: f64,
150 },
151 MultiTaskLearning { task_weights: HashMap<String, f64> },
153 MetaLearning {
155 inner_steps: usize,
156 meta_learning_rate: f64,
157 },
158 ProgressiveTransfer {
160 intermediate_domains: Vec<String>,
161 progression_strategy: ProgressionStrategy,
162 },
163}
164
165#[derive(Debug, Clone, Serialize, Deserialize)]
167pub enum AlignmentMethod {
168 LinearAlignment,
170 NeuralAlignment,
172 AdversarialAlignment,
174 CCA,
176 ProcrustesAlignment,
178 WassersteinAlignment,
180}
181
182#[derive(Debug, Clone, Serialize, Deserialize)]
184pub enum ProgressionStrategy {
185 Sequential,
187 CurriculumBased,
189 SimilarityGuided,
191}
192
193#[derive(Debug, Clone, Serialize, Deserialize)]
195pub enum TransferMetric {
196 TransferAccuracy,
198 AdaptationQuality,
200 EntityAlignmentQuality,
202 SemanticPreservation,
204 StructuralPreservation,
206 TransferEfficiency,
208 CatastrophicForgetting,
210 CrossDomainCoherence,
212 KnowledgeRetention,
214 AdaptationSpeed,
216 TransferRobustness,
218 SemanticDriftDetection,
220 GeneralizationAbility,
222}
223
224#[derive(Debug, Clone, Serialize, Deserialize)]
226pub struct TransferEvaluationResults {
227 pub source_domain: String,
229 pub target_domain: String,
231 pub strategy: TransferStrategy,
233 pub metric_scores: HashMap<String, f64>,
235 pub overall_quality: f64,
237 pub domain_similarity: f64,
239 pub improvement_over_baseline: f64,
241 pub transfer_time: f64,
243 pub detailed_analysis: TransferAnalysis,
245}
246
247#[derive(Debug, Clone, Serialize, Deserialize)]
249pub struct TransferAnalysis {
250 pub entity_alignments: Vec<EntityAlignment>,
252 pub relation_alignments: Vec<RelationAlignment>,
254 pub semantic_shifts: Vec<SemanticShift>,
256 pub structural_changes: StructuralChanges,
258 pub recommendations: Vec<String>,
260}
261
262#[derive(Debug, Clone, Serialize, Deserialize)]
264pub struct EntityAlignment {
265 pub source_entity: String,
267 pub target_entity: String,
269 pub confidence: f64,
271 pub similarity: f64,
273 pub method: String,
275}
276
277#[derive(Debug, Clone, Serialize, Deserialize)]
279pub struct RelationAlignment {
280 pub source_relation: String,
282 pub target_relation: String,
284 pub confidence: f64,
286 pub semantic_similarity: f64,
288 pub structural_similarity: f64,
290}
291
292#[derive(Debug, Clone, Serialize, Deserialize)]
294pub struct SemanticShift {
295 pub concept: String,
297 pub source_meaning: String,
299 pub target_meaning: String,
301 pub shift_magnitude: f64,
303 pub impact: f64,
305}
306
307#[derive(Debug, Clone, Serialize, Deserialize)]
309pub struct StructuralChanges {
310 pub degree_distribution_shift: f64,
312 pub clustering_changes: f64,
314 pub path_length_changes: f64,
316 pub community_structure_changes: f64,
318}
319
320impl CrossDomainTransferManager {
321 pub fn new(config: TransferConfig) -> Self {
323 Self {
324 source_domains: HashMap::new(),
325 target_domains: HashMap::new(),
326 transfer_strategies: vec![
327 TransferStrategy::DirectTransfer,
328 TransferStrategy::FineTuning {
329 learning_rate: 0.001,
330 epochs: 50,
331 freeze_layers: vec![],
332 },
333 TransferStrategy::DomainAdaptation {
334 alignment_method: AlignmentMethod::AdversarialAlignment,
335 regularization_strength: 0.1,
336 },
337 ],
338 transfer_metrics: vec![
339 TransferMetric::TransferAccuracy,
340 TransferMetric::AdaptationQuality,
341 TransferMetric::SemanticPreservation,
342 TransferMetric::StructuralPreservation,
343 ],
344 config,
345 }
346 }
347
348 pub fn register_source_domain(
350 &mut self,
351 domain_id: String,
352 model: Box<dyn EmbeddingModel + Send + Sync>,
353 characteristics: DomainCharacteristics,
354 ) -> Result<()> {
355 let domain_model = DomainModel {
356 domain_id: domain_id.clone(),
357 model,
358 characteristics,
359 entity_mappings: HashMap::new(),
360 vocabulary: HashSet::new(),
361 };
362
363 self.source_domains.insert(domain_id, domain_model);
364 Ok(())
365 }
366
367 pub fn register_target_domain(&mut self, domain_spec: DomainSpecification) -> Result<()> {
369 self.target_domains
370 .insert(domain_spec.domain_id.clone(), domain_spec);
371 Ok(())
372 }
373
374 pub async fn evaluate_transfer(
376 &self,
377 source_domain_id: &str,
378 target_domain_id: &str,
379 strategy: TransferStrategy,
380 ) -> Result<TransferEvaluationResults> {
381 let source_domain = self
382 .source_domains
383 .get(source_domain_id)
384 .ok_or_else(|| anyhow!("Source domain not found: {}", source_domain_id))?;
385
386 let target_domain = self
387 .target_domains
388 .get(target_domain_id)
389 .ok_or_else(|| anyhow!("Target domain not found: {}", target_domain_id))?;
390
391 let start_time = std::time::Instant::now();
392
393 let domain_similarity = self.calculate_domain_similarity(
395 &source_domain.characteristics,
396 &target_domain.characteristics,
397 )?;
398
399 let entity_alignments = self.align_entities(source_domain, target_domain).await?;
401
402 let relation_alignments = self.align_relations(source_domain, target_domain).await?;
404
405 let semantic_shifts = self
407 .analyze_semantic_shifts(source_domain, target_domain)
408 .await?;
409
410 let structural_changes = self.analyze_structural_changes(
412 &source_domain.characteristics,
413 &target_domain.characteristics,
414 )?;
415
416 let mut metric_scores = HashMap::new();
418 for metric in &self.transfer_metrics {
419 let score = self
420 .evaluate_transfer_metric(
421 metric,
422 source_domain,
423 target_domain,
424 &entity_alignments,
425 &relation_alignments,
426 )
427 .await?;
428 metric_scores.insert(format!("{metric:?}"), score);
429 }
430
431 let overall_quality = if metric_scores.is_empty() {
433 0.5 } else {
435 let avg_quality = metric_scores.values().sum::<f64>() / metric_scores.len() as f64;
436 avg_quality.max(0.0) };
438
439 let baseline_performance = 0.1; let improvement_over_baseline = overall_quality - baseline_performance;
442
443 let transfer_time = start_time.elapsed().as_secs_f64();
444
445 let recommendations = self.generate_transfer_recommendations(
447 domain_similarity,
448 &entity_alignments,
449 &semantic_shifts,
450 );
451
452 let detailed_analysis = TransferAnalysis {
453 entity_alignments,
454 relation_alignments,
455 semantic_shifts,
456 structural_changes,
457 recommendations,
458 };
459
460 Ok(TransferEvaluationResults {
461 source_domain: source_domain_id.to_string(),
462 target_domain: target_domain_id.to_string(),
463 strategy,
464 metric_scores,
465 overall_quality,
466 domain_similarity,
467 improvement_over_baseline,
468 transfer_time,
469 detailed_analysis,
470 })
471 }
472
473 pub fn calculate_domain_similarity(
475 &self,
476 source: &DomainCharacteristics,
477 target: &DomainCharacteristics,
478 ) -> Result<f64> {
479 let mut similarity_scores = Vec::new();
480
481 let language_similarity = if source.language == target.language {
483 1.0
484 } else {
485 0.5 };
487 similarity_scores.push(language_similarity);
488
489 let source_entity_types: HashSet<_> = source.entity_types.iter().collect();
491 let target_entity_types: HashSet<_> = target.entity_types.iter().collect();
492 let entity_overlap = source_entity_types
493 .intersection(&target_entity_types)
494 .count() as f64;
495 let entity_similarity =
496 entity_overlap / (source_entity_types.len() + target_entity_types.len()) as f64 * 2.0;
497 similarity_scores.push(entity_similarity);
498
499 let source_relation_types: HashSet<_> = source.relation_types.iter().collect();
501 let target_relation_types: HashSet<_> = target.relation_types.iter().collect();
502 let relation_overlap = source_relation_types
503 .intersection(&target_relation_types)
504 .count() as f64;
505 let relation_similarity = relation_overlap
506 / (source_relation_types.len() + target_relation_types.len()) as f64
507 * 2.0;
508 similarity_scores.push(relation_similarity);
509
510 let size_ratio = (target.size_metrics.num_entities as f64
512 / source.size_metrics.num_entities as f64)
513 .min(source.size_metrics.num_entities as f64 / target.size_metrics.num_entities as f64);
514 similarity_scores.push(size_ratio);
515
516 let complexity_diff = (source.complexity_metrics.semantic_diversity
518 - target.complexity_metrics.semantic_diversity)
519 .abs();
520 let complexity_similarity = (1.0 - complexity_diff).max(0.0);
521 similarity_scores.push(complexity_similarity);
522
523 let overall_similarity =
525 similarity_scores.iter().sum::<f64>() / similarity_scores.len() as f64;
526
527 Ok(overall_similarity)
528 }
529
530 async fn align_entities(
532 &self,
533 source: &DomainModel,
534 target: &DomainSpecification,
535 ) -> Result<Vec<EntityAlignment>> {
536 let mut alignments = Vec::new();
537
538 let source_entities = source.model.get_entities();
539 let target_entities = self.extract_entities_from_triples(&target.training_data);
540
541 for source_entity in &source_entities {
543 for target_entity in &target_entities {
544 let similarity = self.calculate_string_similarity(source_entity, target_entity);
545
546 if similarity > 0.7 {
547 alignments.push(EntityAlignment {
549 source_entity: source_entity.clone(),
550 target_entity: target_entity.clone(),
551 confidence: similarity,
552 similarity,
553 method: "string_similarity".to_string(),
554 });
555 }
556 }
557 }
558
559 for source_entity in source_entities.iter().take(50) {
561 if let Ok(source_embedding) = source.model.get_entity_embedding(source_entity) {
563 let mut best_match = None;
564 let mut best_similarity = 0.0;
565
566 for target_entity in target_entities.iter().take(50) {
567 let target_embedding = self.create_simple_embedding(target_entity);
569 let similarity = self.cosine_similarity(&source_embedding, &target_embedding);
570
571 if similarity > best_similarity && similarity > 0.5 {
572 best_similarity = similarity;
573 best_match = Some(target_entity.clone());
574 }
575 }
576
577 if let Some(target_entity) = best_match {
578 alignments.push(EntityAlignment {
579 source_entity: source_entity.clone(),
580 target_entity,
581 confidence: best_similarity,
582 similarity: best_similarity,
583 method: "semantic_embedding".to_string(),
584 });
585 }
586 }
587 }
588
589 Ok(alignments)
590 }
591
592 async fn align_relations(
594 &self,
595 source: &DomainModel,
596 target: &DomainSpecification,
597 ) -> Result<Vec<RelationAlignment>> {
598 let mut alignments = Vec::new();
599
600 let source_relations = source.model.get_relations();
601 let target_relations = self.extract_relations_from_triples(&target.training_data);
602
603 let source_stats = source.model.get_stats();
609 let source_avg_relation_frequency = if source_stats.num_relations > 0 {
610 source_stats.num_triples as f64 / source_stats.num_relations as f64
611 } else {
612 0.0
613 };
614
615 let mut target_relation_frequency: HashMap<&str, usize> = HashMap::new();
616 for (_, predicate, _) in &target.training_data {
617 *target_relation_frequency
618 .entry(predicate.as_str())
619 .or_insert(0) += 1;
620 }
621
622 for source_relation in &source_relations {
623 for target_relation in &target_relations {
624 let semantic_similarity =
625 self.calculate_string_similarity(source_relation, target_relation);
626
627 let target_frequency = target_relation_frequency
631 .get(target_relation.as_str())
632 .copied()
633 .unwrap_or(0) as f64;
634 let max_frequency = source_avg_relation_frequency.max(target_frequency).max(1.0);
635 let structural_similarity =
636 1.0 - (source_avg_relation_frequency - target_frequency).abs() / max_frequency;
637
638 if semantic_similarity > 0.6 {
639 alignments.push(RelationAlignment {
640 source_relation: source_relation.clone(),
641 target_relation: target_relation.clone(),
642 confidence: (semantic_similarity + structural_similarity) / 2.0,
643 semantic_similarity,
644 structural_similarity,
645 });
646 }
647 }
648 }
649
650 Ok(alignments)
651 }
652
653 async fn analyze_semantic_shifts(
655 &self,
656 source: &DomainModel,
657 target: &DomainSpecification,
658 ) -> Result<Vec<SemanticShift>> {
659 let mut shifts = Vec::new();
660
661 let source_entities = source.model.get_entities();
663 let target_entities = self.extract_entities_from_triples(&target.training_data);
664
665 for source_entity in source_entities.iter().take(20) {
666 for target_entity in target_entities.iter().take(20) {
667 if self.calculate_string_similarity(source_entity, target_entity) > 0.8 {
668 let shift_magnitude = self.calculate_semantic_shift_magnitude(
670 source_entity,
671 target_entity,
672 source,
673 target,
674 )?;
675
676 if shift_magnitude > 0.3 {
677 shifts.push(SemanticShift {
678 concept: source_entity.clone(),
679 source_meaning: format!("Source domain context: {source_entity}"),
680 target_meaning: format!("Target domain context: {target_entity}"),
681 shift_magnitude,
682 impact: shift_magnitude * 0.5, });
684 }
685 }
686 }
687 }
688
689 Ok(shifts)
690 }
691
692 fn analyze_structural_changes(
708 &self,
709 source: &DomainCharacteristics,
710 target: &DomainCharacteristics,
711 ) -> Result<StructuralChanges> {
712 let relative_change = |source_value: f64, target_value: f64| -> f64 {
713 let denom = source_value.abs().max(f64::EPSILON);
714 ((source_value - target_value).abs() / denom).min(1.0)
715 };
716
717 let degree_distribution_shift =
719 (source.size_metrics.avg_entity_degree - target.size_metrics.avg_entity_degree).abs()
720 / source.size_metrics.avg_entity_degree;
721
722 let clustering_changes = relative_change(
723 source.size_metrics.graph_density,
724 target.size_metrics.graph_density,
725 );
726 let path_length_changes = relative_change(
727 source.complexity_metrics.hierarchical_depth as f64,
728 target.complexity_metrics.hierarchical_depth as f64,
729 );
730 let community_structure_changes = relative_change(
731 source.complexity_metrics.structural_complexity,
732 target.complexity_metrics.structural_complexity,
733 );
734
735 Ok(StructuralChanges {
736 degree_distribution_shift,
737 clustering_changes,
738 path_length_changes,
739 community_structure_changes,
740 })
741 }
742
743 async fn evaluate_transfer_metric(
745 &self,
746 metric: &TransferMetric,
747 source: &DomainModel,
748 target: &DomainSpecification,
749 entity_alignments: &[EntityAlignment],
750 relation_alignments: &[RelationAlignment],
751 ) -> Result<f64> {
752 match metric {
753 TransferMetric::TransferAccuracy => {
754 self.calculate_transfer_accuracy(source, target).await
756 }
757 TransferMetric::AdaptationQuality => {
758 if entity_alignments.is_empty() {
760 Ok(0.5) } else {
762 Ok(entity_alignments.iter().map(|a| a.confidence).sum::<f64>()
763 / entity_alignments.len() as f64)
764 }
765 }
766 TransferMetric::EntityAlignmentQuality => {
767 if entity_alignments.is_empty() {
769 Ok(0.5) } else {
771 Ok(entity_alignments
772 .iter()
773 .filter(|a| a.confidence > 0.7)
774 .count() as f64
775 / entity_alignments.len() as f64)
776 }
777 }
778 TransferMetric::SemanticPreservation => {
779 self.calculate_semantic_preservation(source, target, entity_alignments)
781 .await
782 }
783 TransferMetric::StructuralPreservation => {
784 self.calculate_structural_preservation(source, target, relation_alignments)
786 .await
787 }
788 TransferMetric::TransferEfficiency => {
789 self.calculate_transfer_efficiency(source, target).await
791 }
792 TransferMetric::CatastrophicForgetting => {
793 self.calculate_catastrophic_forgetting(source, target).await
795 }
796 TransferMetric::CrossDomainCoherence => {
797 self.calculate_cross_domain_coherence(source, target, entity_alignments)
799 .await
800 }
801 TransferMetric::KnowledgeRetention => {
802 self.calculate_knowledge_retention(source, target).await
804 }
805 TransferMetric::AdaptationSpeed => {
806 self.calculate_adaptation_speed(source, target).await
808 }
809 TransferMetric::TransferRobustness => {
810 self.calculate_transfer_robustness(source, target).await
812 }
813 TransferMetric::SemanticDriftDetection => {
814 self.calculate_semantic_drift_detection(source, target)
816 .await
817 }
818 TransferMetric::GeneralizationAbility => {
819 self.calculate_generalization_ability(source, target).await
821 }
822 }
823 }
824
825 async fn calculate_transfer_accuracy(
827 &self,
828 source: &DomainModel,
829 target: &DomainSpecification,
830 ) -> Result<f64> {
831 let mut correct_predictions = 0;
832 let total_predictions = target
833 .test_data
834 .len()
835 .min(self.config.evaluation_sample_size);
836
837 if total_predictions == 0 {
838 return Ok(0.5); }
840
841 for (subject, predicate, object) in target.test_data.iter().take(total_predictions) {
842 if let Ok(score) = source.model.score_triple(subject, predicate, object) {
844 if score > 0.0 {
846 correct_predictions += 1;
847 }
848 }
849 }
850
851 Ok(correct_predictions as f64 / total_predictions as f64)
852 }
853
854 async fn calculate_semantic_preservation(
856 &self,
857 source: &DomainModel,
858 _target: &DomainSpecification,
859 entity_alignments: &[EntityAlignment],
860 ) -> Result<f64> {
861 if entity_alignments.is_empty() {
862 return Ok(0.0);
863 }
864
865 let mut preservation_scores = Vec::new();
866
867 for alignment in entity_alignments.iter().take(20) {
868 if let Ok(source_embedding) =
870 source.model.get_entity_embedding(&alignment.source_entity)
871 {
872 let target_embedding = self.create_simple_embedding(&alignment.target_entity);
874
875 let preservation = self.cosine_similarity(&source_embedding, &target_embedding);
877 preservation_scores.push(preservation);
878 }
879 }
880
881 if preservation_scores.is_empty() {
882 Ok(0.0)
883 } else {
884 Ok(preservation_scores.iter().sum::<f64>() / preservation_scores.len() as f64)
885 }
886 }
887
888 async fn calculate_structural_preservation(
890 &self,
891 _source: &DomainModel,
892 _target: &DomainSpecification,
893 relation_alignments: &[RelationAlignment],
894 ) -> Result<f64> {
895 if relation_alignments.is_empty() {
896 return Ok(0.5); }
898
899 let avg_structural_similarity = relation_alignments
901 .iter()
902 .map(|a| a.structural_similarity)
903 .sum::<f64>()
904 / relation_alignments.len() as f64;
905
906 Ok(avg_structural_similarity)
907 }
908
909 fn generate_transfer_recommendations(
911 &self,
912 domain_similarity: f64,
913 entity_alignments: &[EntityAlignment],
914 semantic_shifts: &[SemanticShift],
915 ) -> Vec<String> {
916 let mut recommendations = Vec::new();
917
918 if domain_similarity < 0.3 {
919 recommendations.push(
920 "Low domain similarity detected. Consider using domain adaptation techniques."
921 .to_string(),
922 );
923 }
924
925 if entity_alignments.len() < 10 {
926 recommendations.push(
927 "Few entity alignments found. Consider improving entity linking methods."
928 .to_string(),
929 );
930 }
931
932 let high_shift_count = semantic_shifts
933 .iter()
934 .filter(|s| s.shift_magnitude > 0.5)
935 .count();
936 if high_shift_count > 5 {
937 recommendations.push(
938 "Significant semantic shifts detected. Consider gradual domain adaptation."
939 .to_string(),
940 );
941 }
942
943 if domain_similarity > 0.7 {
944 recommendations
945 .push("High domain similarity. Direct transfer should work well.".to_string());
946 }
947
948 recommendations
949 }
950
951 fn extract_entities_from_triples(
953 &self,
954 triples: &[(String, String, String)],
955 ) -> HashSet<String> {
956 let mut entities = HashSet::new();
957 for (subject, _, object) in triples {
958 entities.insert(subject.clone());
959 entities.insert(object.clone());
960 }
961 entities.into_iter().collect::<HashSet<_>>()
962 }
963
964 fn extract_relations_from_triples(
966 &self,
967 triples: &[(String, String, String)],
968 ) -> HashSet<String> {
969 triples
970 .iter()
971 .map(|(_, predicate, _)| predicate.clone())
972 .collect()
973 }
974
975 fn calculate_string_similarity(&self, s1: &str, s2: &str) -> f64 {
977 if s1 == s2 {
978 return 1.0;
979 }
980
981 let n = 3;
983 let ngrams1: HashSet<String> = s1
984 .chars()
985 .collect::<Vec<_>>()
986 .windows(n)
987 .map(|w| w.iter().collect())
988 .collect();
989 let ngrams2: HashSet<String> = s2
990 .chars()
991 .collect::<Vec<_>>()
992 .windows(n)
993 .map(|w| w.iter().collect())
994 .collect();
995
996 if ngrams1.is_empty() && ngrams2.is_empty() {
997 return 1.0;
998 }
999
1000 let intersection = ngrams1.intersection(&ngrams2).count();
1001 let union = ngrams1.union(&ngrams2).count();
1002
1003 intersection as f64 / union as f64
1004 }
1005
1006 fn create_simple_embedding(&self, entity: &str) -> Vector {
1008 let mut embedding = vec![0.0f32; 100]; for (i, byte) in entity.bytes().enumerate() {
1011 if i >= embedding.len() {
1012 break;
1013 }
1014 embedding[i] = (byte as f32) / 255.0;
1015 }
1016 Vector::new(embedding)
1017 }
1018
1019 fn cosine_similarity(&self, v1: &Vector, v2: &Vector) -> f64 {
1021 let dot_product: f32 = v1
1022 .values
1023 .iter()
1024 .zip(v2.values.iter())
1025 .map(|(a, b)| a * b)
1026 .sum();
1027 let norm_a: f32 = v1.values.iter().map(|x| x * x).sum::<f32>().sqrt();
1028 let norm_b: f32 = v2.values.iter().map(|x| x * x).sum::<f32>().sqrt();
1029
1030 if norm_a > 0.0 && norm_b > 0.0 {
1031 (dot_product / (norm_a * norm_b)) as f64
1032 } else {
1033 0.0
1034 }
1035 }
1036
1037 fn calculate_semantic_shift_magnitude(
1039 &self,
1040 _source_entity: &str,
1041 _target_entity: &str,
1042 _source: &DomainModel,
1043 _target: &DomainSpecification,
1044 ) -> Result<f64> {
1045 Ok({
1047 let mut random = Random::default();
1048 random.random::<f64>() * 0.8
1049 }) }
1051
1052 pub fn get_source_domains(&self) -> Vec<String> {
1054 self.source_domains.keys().cloned().collect()
1055 }
1056
1057 pub fn get_target_domains(&self) -> Vec<String> {
1059 self.target_domains.keys().cloned().collect()
1060 }
1061
1062 pub fn get_domain_characteristics(&self, domain_id: &str) -> Option<&DomainCharacteristics> {
1064 self.source_domains
1065 .get(domain_id)
1066 .map(|d| &d.characteristics)
1067 .or_else(|| {
1068 self.target_domains
1069 .get(domain_id)
1070 .map(|d| &d.characteristics)
1071 })
1072 }
1073
1074 async fn calculate_transfer_efficiency(
1076 &self,
1077 source: &DomainModel,
1078 target: &DomainSpecification,
1079 ) -> Result<f64> {
1080 let start_time = std::time::Instant::now();
1081
1082 let domain_similarity =
1084 self.calculate_domain_similarity(&source.characteristics, &target.characteristics)?;
1085
1086 let transfer_accuracy = self.calculate_transfer_accuracy(source, target).await?;
1088 let transfer_time = start_time.elapsed().as_secs_f64();
1089
1090 let normalized_time = (transfer_time / 60.0).clamp(0.01, 1.0); let efficiency = (transfer_accuracy * domain_similarity) / normalized_time;
1093
1094 Ok(efficiency.clamp(0.0, 1.0))
1095 }
1096
1097 async fn calculate_catastrophic_forgetting(
1099 &self,
1100 source: &DomainModel,
1101 target: &DomainSpecification,
1102 ) -> Result<f64> {
1103 let source_entities = source.model.get_entities();
1105 let sample_size = source_entities.len().min(20);
1106
1107 if sample_size == 0 {
1108 return Ok(0.0);
1109 }
1110
1111 let mut forgetting_scores = Vec::new();
1112
1113 for entity in source_entities.iter().take(sample_size) {
1115 if let Ok(_source_embedding) = source.model.get_entity_embedding(entity) {
1116 let target_entities = self.extract_entities_from_triples(&target.training_data);
1118 let domain_overlap = target_entities.contains(entity);
1119
1120 let degradation = if domain_overlap {
1121 let mut random = Random::default();
1123 0.1 + random.random::<f64>() * 0.2
1124 } else {
1125 let mut random = Random::default();
1127 0.3 + random.random::<f64>() * 0.4
1128 };
1129
1130 forgetting_scores.push(degradation);
1131 }
1132 }
1133
1134 if forgetting_scores.is_empty() {
1135 Ok(0.1) } else {
1137 let avg_forgetting =
1138 forgetting_scores.iter().sum::<f64>() / forgetting_scores.len() as f64;
1139 Ok(avg_forgetting.clamp(0.0, 1.0))
1140 }
1141 }
1142
1143 async fn calculate_cross_domain_coherence(
1145 &self,
1146 source: &DomainModel,
1147 target: &DomainSpecification,
1148 entity_alignments: &[EntityAlignment],
1149 ) -> Result<f64> {
1150 if entity_alignments.is_empty() {
1151 return Ok(0.5);
1152 }
1153
1154 let mut coherence_scores = Vec::new();
1155
1156 for alignment in entity_alignments.iter().take(15) {
1158 if alignment.confidence > 0.6 {
1159 if let Ok(source_embedding) =
1160 source.model.get_entity_embedding(&alignment.source_entity)
1161 {
1162 let target_embedding = self.create_simple_embedding(&alignment.target_entity);
1163
1164 let embedding_coherence =
1166 self.cosine_similarity(&source_embedding, &target_embedding);
1167
1168 let source_neighbors =
1170 self.get_source_neighbors(&alignment.source_entity, source);
1171 let target_neighbors =
1172 self.get_target_neighbors(&alignment.target_entity, target);
1173 let neighborhood_coherence = self
1174 .calculate_neighborhood_similarity(&source_neighbors, &target_neighbors);
1175
1176 let combined_coherence = (embedding_coherence + neighborhood_coherence) / 2.0;
1178 coherence_scores.push(combined_coherence);
1179 }
1180 }
1181 }
1182
1183 if coherence_scores.is_empty() {
1184 Ok(0.5)
1185 } else {
1186 let avg_coherence =
1187 coherence_scores.iter().sum::<f64>() / coherence_scores.len() as f64;
1188 Ok(avg_coherence.clamp(0.0, 1.0))
1189 }
1190 }
1191
1192 async fn calculate_knowledge_retention(
1194 &self,
1195 _source: &DomainModel,
1196 _target: &DomainSpecification,
1197 ) -> Result<f64> {
1198 Ok(0.85)
1200 }
1201
1202 async fn calculate_adaptation_speed(
1204 &self,
1205 _source: &DomainModel,
1206 _target: &DomainSpecification,
1207 ) -> Result<f64> {
1208 Ok(0.75)
1210 }
1211
1212 async fn calculate_transfer_robustness(
1214 &self,
1215 _source: &DomainModel,
1216 _target: &DomainSpecification,
1217 ) -> Result<f64> {
1218 Ok(0.8)
1220 }
1221
1222 async fn calculate_semantic_drift_detection(
1224 &self,
1225 _source: &DomainModel,
1226 _target: &DomainSpecification,
1227 ) -> Result<f64> {
1228 Ok(0.7)
1230 }
1231
1232 async fn calculate_generalization_ability(
1234 &self,
1235 _source: &DomainModel,
1236 _target: &DomainSpecification,
1237 ) -> Result<f64> {
1238 Ok(0.8)
1240 }
1241
1242 fn get_source_neighbors(&self, _entity: &str, source: &DomainModel) -> Vec<String> {
1244 let relations = source.model.get_relations();
1246 relations.into_iter().take(5).collect()
1247 }
1248
1249 fn get_target_neighbors(&self, entity: &str, target: &DomainSpecification) -> Vec<String> {
1251 let mut neighbors = Vec::new();
1252 for (subject, predicate, object) in &target.training_data {
1253 if subject == entity {
1254 neighbors.push(object.clone());
1255 neighbors.push(predicate.clone());
1256 } else if object == entity {
1257 neighbors.push(subject.clone());
1258 neighbors.push(predicate.clone());
1259 }
1260 }
1261 neighbors.into_iter().take(5).collect()
1262 }
1263
1264 fn calculate_neighborhood_similarity(
1266 &self,
1267 source_neighbors: &[String],
1268 target_neighbors: &[String],
1269 ) -> f64 {
1270 if source_neighbors.is_empty() && target_neighbors.is_empty() {
1271 return 1.0;
1272 }
1273
1274 if source_neighbors.is_empty() || target_neighbors.is_empty() {
1275 return 0.0;
1276 }
1277
1278 let source_set: HashSet<&String> = source_neighbors.iter().collect();
1279 let target_set: HashSet<&String> = target_neighbors.iter().collect();
1280
1281 let intersection = source_set.intersection(&target_set).count();
1282 let union = source_set.union(&target_set).count();
1283
1284 if union == 0 {
1285 0.0
1286 } else {
1287 intersection as f64 / union as f64
1288 }
1289 }
1290}
1291
1292pub struct TransferUtils;
1294
1295impl TransferUtils {
1296 pub fn analyze_domain_from_triples(
1298 _domain_id: String,
1299 triples: &[(String, String, String)],
1300 ) -> DomainCharacteristics {
1301 let mut entities = HashSet::new();
1302 let mut relations = HashSet::new();
1303
1304 for (subject, predicate, object) in triples {
1305 entities.insert(subject.clone());
1306 entities.insert(object.clone());
1307 relations.insert(predicate.clone());
1308 }
1309
1310 let num_entities = entities.len();
1311 let num_relations = relations.len();
1312 let num_triples = triples.len();
1313
1314 let mut entity_degrees = HashMap::new();
1316 for (subject, _, object) in triples {
1317 *entity_degrees.entry(subject.clone()).or_insert(0) += 1;
1318 *entity_degrees.entry(object.clone()).or_insert(0) += 1;
1319 }
1320 let avg_entity_degree = if num_entities > 0 {
1321 entity_degrees.values().sum::<usize>() as f64 / num_entities as f64
1322 } else {
1323 0.0
1324 };
1325
1326 let max_possible_edges = num_entities * (num_entities - 1);
1328 let graph_density = if max_possible_edges > 0 {
1329 num_triples as f64 / max_possible_edges as f64
1330 } else {
1331 0.0
1332 };
1333
1334 DomainCharacteristics {
1335 domain_type: "unknown".to_string(),
1336 language: "unknown".to_string(),
1337 entity_types: vec!["Entity".to_string()], relation_types: relations.into_iter().collect(),
1339 size_metrics: DomainSizeMetrics {
1340 num_entities,
1341 num_relations,
1342 num_triples,
1343 avg_entity_degree,
1344 graph_density,
1345 },
1346 complexity_metrics: DomainComplexityMetrics {
1347 entity_type_diversity: 1, relation_type_diversity: num_relations,
1349 hierarchical_depth: 3, semantic_diversity: 0.5, structural_complexity: avg_entity_degree / 10.0, },
1353 }
1354 }
1355
1356 pub fn create_test_domain_specification(
1358 domain_id: String,
1359 training_data: Vec<(String, String, String)>,
1360 ) -> DomainSpecification {
1361 let total = training_data.len();
1363 let train_size = (total as f64 * 0.7) as usize;
1364 let val_size = (total as f64 * 0.15) as usize;
1365
1366 let training = training_data[..train_size].to_vec();
1367 let validation = training_data[train_size..train_size + val_size].to_vec();
1368 let test = training_data[train_size + val_size..].to_vec();
1369
1370 let characteristics = Self::analyze_domain_from_triples(domain_id.clone(), &training);
1371
1372 DomainSpecification {
1373 domain_id,
1374 characteristics,
1375 training_data: training,
1376 validation_data: validation,
1377 test_data: test,
1378 }
1379 }
1380}
1381
1382#[cfg(test)]
1383mod tests {
1384 use super::*;
1385 use crate::models::transe::TransE;
1386
1387 #[test]
1388 fn test_transfer_config_default() {
1389 let config = TransferConfig::default();
1390 assert!(config.enable_domain_adaptation);
1391 assert!(config.use_adversarial_alignment);
1392 assert_eq!(config.max_alignment_iterations, 100);
1393 }
1394
1395 #[test]
1396 fn test_domain_characteristics_creation() {
1397 let triples = vec![
1398 ("alice".to_string(), "knows".to_string(), "bob".to_string()),
1399 ("bob".to_string(), "likes".to_string(), "pizza".to_string()),
1400 (
1401 "alice".to_string(),
1402 "likes".to_string(),
1403 "coffee".to_string(),
1404 ),
1405 ];
1406
1407 let characteristics =
1408 TransferUtils::analyze_domain_from_triples("test_domain".to_string(), &triples);
1409
1410 assert_eq!(characteristics.size_metrics.num_triples, 3);
1411 assert_eq!(characteristics.size_metrics.num_entities, 4); assert_eq!(characteristics.size_metrics.num_relations, 2); }
1414
1415 #[test]
1416 fn test_string_similarity() {
1417 let manager = CrossDomainTransferManager::new(TransferConfig::default());
1418
1419 let sim1 = manager.calculate_string_similarity("hello", "hello");
1420 assert_eq!(sim1, 1.0);
1421
1422 let sim2 = manager.calculate_string_similarity("hello", "world");
1423 assert!(sim2 < 0.5);
1424
1425 let sim3 = manager.calculate_string_similarity("testing", "test");
1426 assert!(sim3 > 0.3);
1427 }
1428
1429 #[tokio::test]
1430 async fn test_transfer_evaluation() {
1431 let mut manager = CrossDomainTransferManager::new(TransferConfig::default());
1432
1433 let source_model = Box::new(TransE::new(Default::default()));
1435 let source_characteristics = DomainCharacteristics {
1436 domain_type: "test".to_string(),
1437 language: "en".to_string(),
1438 entity_types: vec!["Person".to_string()],
1439 relation_types: vec!["knows".to_string()],
1440 size_metrics: DomainSizeMetrics {
1441 num_entities: 100,
1442 num_relations: 10,
1443 num_triples: 500,
1444 avg_entity_degree: 5.0,
1445 graph_density: 0.01,
1446 },
1447 complexity_metrics: DomainComplexityMetrics {
1448 entity_type_diversity: 2,
1449 relation_type_diversity: 10,
1450 hierarchical_depth: 3,
1451 semantic_diversity: 0.6,
1452 structural_complexity: 0.5,
1453 },
1454 };
1455
1456 manager
1457 .register_source_domain("source".to_string(), source_model, source_characteristics)
1458 .expect("should succeed");
1459
1460 let target_spec = TransferUtils::create_test_domain_specification(
1462 "target".to_string(),
1463 vec![
1464 ("alice".to_string(), "knows".to_string(), "bob".to_string()),
1465 (
1466 "bob".to_string(),
1467 "knows".to_string(),
1468 "charlie".to_string(),
1469 ),
1470 ],
1471 );
1472
1473 manager
1474 .register_target_domain(target_spec)
1475 .expect("should succeed");
1476
1477 let results = manager
1479 .evaluate_transfer("source", "target", TransferStrategy::DirectTransfer)
1480 .await;
1481
1482 assert!(results.is_ok());
1483 let results = results.expect("should succeed");
1484 assert_eq!(results.source_domain, "source");
1485 assert_eq!(results.target_domain, "target");
1486 assert!(results.overall_quality >= 0.0);
1487 assert!(results.overall_quality <= 1.0);
1488 }
1489
1490 #[test]
1491 fn test_domain_similarity_calculation() {
1492 let manager = CrossDomainTransferManager::new(TransferConfig::default());
1493
1494 let source = DomainCharacteristics {
1495 domain_type: "biomedical".to_string(),
1496 language: "en".to_string(),
1497 entity_types: vec!["Gene".to_string(), "Disease".to_string()],
1498 relation_types: vec!["causes".to_string(), "treats".to_string()],
1499 size_metrics: DomainSizeMetrics {
1500 num_entities: 1000,
1501 num_relations: 50,
1502 num_triples: 5000,
1503 avg_entity_degree: 5.0,
1504 graph_density: 0.005,
1505 },
1506 complexity_metrics: DomainComplexityMetrics {
1507 entity_type_diversity: 2,
1508 relation_type_diversity: 50,
1509 hierarchical_depth: 4,
1510 semantic_diversity: 0.7,
1511 structural_complexity: 0.6,
1512 },
1513 };
1514
1515 let target = DomainCharacteristics {
1516 domain_type: "medical".to_string(),
1517 language: "en".to_string(),
1518 entity_types: vec!["Gene".to_string(), "Drug".to_string()],
1519 relation_types: vec!["treats".to_string(), "interacts".to_string()],
1520 size_metrics: DomainSizeMetrics {
1521 num_entities: 800,
1522 num_relations: 40,
1523 num_triples: 4000,
1524 avg_entity_degree: 5.0,
1525 graph_density: 0.006,
1526 },
1527 complexity_metrics: DomainComplexityMetrics {
1528 entity_type_diversity: 2,
1529 relation_type_diversity: 40,
1530 hierarchical_depth: 3,
1531 semantic_diversity: 0.6,
1532 structural_complexity: 0.5,
1533 },
1534 };
1535
1536 let similarity = manager
1537 .calculate_domain_similarity(&source, &target)
1538 .expect("should succeed");
1539 assert!(similarity > 0.0);
1540 assert!(similarity <= 1.0);
1541
1542 assert!(similarity > 0.2);
1544 }
1545
1546 #[test]
1550 fn test_analyze_structural_changes_varies_with_real_data() {
1551 let manager = CrossDomainTransferManager::new(TransferConfig::default());
1552
1553 let base_characteristics =
1554 |graph_density: f64, hierarchical_depth: usize, structural_complexity: f64| {
1555 DomainCharacteristics {
1556 domain_type: "test".to_string(),
1557 language: "en".to_string(),
1558 entity_types: vec![],
1559 relation_types: vec![],
1560 size_metrics: DomainSizeMetrics {
1561 num_entities: 100,
1562 num_relations: 10,
1563 num_triples: 500,
1564 avg_entity_degree: 5.0,
1565 graph_density,
1566 },
1567 complexity_metrics: DomainComplexityMetrics {
1568 entity_type_diversity: 2,
1569 relation_type_diversity: 2,
1570 hierarchical_depth,
1571 semantic_diversity: 0.5,
1572 structural_complexity,
1573 },
1574 }
1575 };
1576
1577 let source = base_characteristics(0.1, 4, 0.5);
1578
1579 let identical_target = base_characteristics(0.1, 4, 0.5);
1581 let no_change = manager
1582 .analyze_structural_changes(&source, &identical_target)
1583 .expect("should succeed");
1584 assert_eq!(no_change.clustering_changes, 0.0);
1585 assert_eq!(no_change.path_length_changes, 0.0);
1586 assert_eq!(no_change.community_structure_changes, 0.0);
1587
1588 let different_target = base_characteristics(0.5, 8, 0.9);
1591 let changed = manager
1592 .analyze_structural_changes(&source, &different_target)
1593 .expect("should succeed");
1594 assert!(
1595 changed.clustering_changes > 0.0,
1596 "clustering_changes = {}",
1597 changed.clustering_changes
1598 );
1599 assert!(
1600 changed.path_length_changes > 0.0,
1601 "path_length_changes = {}",
1602 changed.path_length_changes
1603 );
1604 assert!(
1605 changed.community_structure_changes > 0.0,
1606 "community_structure_changes = {}",
1607 changed.community_structure_changes
1608 );
1609 }
1610}