1use crate::adaptive_selection::{OptimizerType, ProblemCharacteristics};
8use crate::error::{OptimError, Result};
9use scirs2_core::ndarray::ScalarOperand;
10use scirs2_core::numeric::{Float, ToPrimitive};
11use std::collections::HashMap;
12use std::fmt::Debug;
13
14fn cast_scalar<A: Float, T: ToPrimitive>(value: T) -> Result<A> {
22 A::from(value).ok_or_else(|| {
23 OptimError::InvalidConfig(
24 "failed to convert a numeric value to the domain selector's scalar type".to_string(),
25 )
26 })
27}
28
29#[derive(Debug, Clone)]
31pub enum DomainStrategy {
32 ComputerVision {
34 resolution_adaptive: bool,
36 batch_norm_tuning: bool,
38 augmentation_aware: bool,
40 },
41 NaturalLanguage {
43 sequence_adaptive: bool,
45 attention_optimized: bool,
47 vocab_aware: bool,
49 },
50 RecommendationSystems {
52 collaborative_filtering: bool,
54 matrix_factorization: bool,
56 cold_start_aware: bool,
58 },
59 TimeSeries {
61 temporal_aware: bool,
63 seasonality_adaptive: bool,
65 multi_step: bool,
67 },
68 ReinforcementLearning {
70 policy_gradient: bool,
72 value_function: bool,
74 exploration_aware: bool,
76 },
77 ScientificComputing {
79 stability_focused: bool,
81 precision_critical: bool,
83 sparse_optimized: bool,
85 },
86}
87
88#[derive(Debug, Clone)]
90pub struct DomainConfig<A: Float> {
91 pub base_learning_rate: A,
93 pub recommended_batch_sizes: Vec<usize>,
95 pub gradient_clip_values: Vec<A>,
97 pub regularization_range: (A, A),
99 pub optimizer_ranking: Vec<OptimizerType>,
101 pub domain_params: HashMap<String, A>,
103}
104
105#[derive(Debug)]
107pub struct DomainSpecificSelector<A: Float> {
108 strategy: DomainStrategy,
110 config: DomainConfig<A>,
112 domain_performance: HashMap<String, Vec<DomainPerformanceMetrics<A>>>,
114 transfer_knowledge: Vec<CrossDomainKnowledge<A>>,
116 currentcontext: Option<OptimizationContext<A>>,
118}
119
120#[derive(Debug, Clone)]
122pub struct DomainPerformanceMetrics<A: Float> {
123 pub validation_accuracy: A,
125 pub domain_specific_score: A,
127 pub stability_score: A,
129 pub convergence_epochs: usize,
131 pub resource_efficiency: A,
133 pub transfer_score: A,
135}
136
137#[derive(Debug, Clone)]
139pub struct CrossDomainKnowledge<A: Float> {
140 pub source_domain: String,
142 pub target_domain: String,
144 pub transferable_params: HashMap<String, A>,
146 pub transfer_score: A,
148 pub successful_strategy: OptimizerType,
150}
151
152#[derive(Debug, Clone)]
154pub struct OptimizationContext<A: Float> {
155 pub problem_chars: ProblemCharacteristics,
157 pub resource_constraints: ResourceConstraints<A>,
159 pub training_config: TrainingConfiguration<A>,
161 pub domain_metadata: HashMap<String, String>,
163}
164
165#[derive(Debug, Clone)]
167pub struct ResourceConstraints<A: Float> {
168 pub max_memory: usize,
170 pub max_time: A,
172 pub gpu_available: bool,
174 pub distributed_capable: bool,
176 pub energy_efficient: bool,
178}
179
180#[derive(Debug, Clone)]
182pub struct TrainingConfiguration<A: Float> {
183 pub max_epochs: usize,
185 pub early_stopping_patience: usize,
187 pub validation_frequency: usize,
189 pub lr_schedule_type: LearningRateScheduleType,
191 pub regularization_approach: RegularizationApproach<A>,
193}
194
195#[derive(Debug, Clone)]
197pub enum LearningRateScheduleType {
198 Constant,
200 ExponentialDecay {
202 decay_rate: f64,
204 },
205 CosineAnnealing {
207 t_max: usize,
209 },
210 ReduceOnPlateau {
212 patience: usize,
214 factor: f64,
216 },
217 OneCycle {
219 max_lr: f64,
221 },
222}
223
224#[derive(Debug, Clone)]
226pub enum RegularizationApproach<A: Float> {
227 L2Only {
229 weight: A,
231 },
232 L1Only {
234 weight: A,
236 },
237 ElasticNet {
239 l1_weight: A,
241 l2_weight: A,
243 },
244 Dropout {
246 dropout_rate: A,
248 },
249 Combined {
251 l2_weight: A,
253 dropout_rate: A,
255 additional_techniques: Vec<String>,
257 },
258}
259
260impl<A: Float + ScalarOperand + Debug + std::iter::Sum + Send + Sync> DomainSpecificSelector<A> {
261 pub fn new(strategy: DomainStrategy) -> Self {
263 let config = Self::default_config_for_strategy(&strategy);
264
265 Self {
266 strategy,
267 config,
268 domain_performance: HashMap::new(),
269 transfer_knowledge: Vec::new(),
270 currentcontext: None,
271 }
272 }
273
274 pub fn setcontext(&mut self, context: OptimizationContext<A>) {
276 self.currentcontext = Some(context);
277 }
278
279 pub fn select_optimal_config(&mut self) -> Result<DomainOptimizationConfig<A>> {
281 let context = self
282 .currentcontext
283 .as_ref()
284 .ok_or_else(|| OptimError::InvalidConfig("No optimization context set".to_string()))?;
285
286 match &self.strategy {
287 DomainStrategy::ComputerVision {
288 resolution_adaptive,
289 batch_norm_tuning,
290 augmentation_aware,
291 } => self.optimize_computer_vision(
292 context,
293 *resolution_adaptive,
294 *batch_norm_tuning,
295 *augmentation_aware,
296 ),
297 DomainStrategy::NaturalLanguage {
298 sequence_adaptive,
299 attention_optimized,
300 vocab_aware,
301 } => self.optimize_natural_language(
302 context,
303 *sequence_adaptive,
304 *attention_optimized,
305 *vocab_aware,
306 ),
307 DomainStrategy::RecommendationSystems {
308 collaborative_filtering,
309 matrix_factorization,
310 cold_start_aware,
311 } => self.optimize_recommendation_systems(
312 context,
313 *collaborative_filtering,
314 *matrix_factorization,
315 *cold_start_aware,
316 ),
317 DomainStrategy::TimeSeries {
318 temporal_aware,
319 seasonality_adaptive,
320 multi_step,
321 } => self.optimize_time_series(
322 context,
323 *temporal_aware,
324 *seasonality_adaptive,
325 *multi_step,
326 ),
327 DomainStrategy::ReinforcementLearning {
328 policy_gradient,
329 value_function,
330 exploration_aware,
331 } => self.optimize_reinforcement_learning(
332 context,
333 *policy_gradient,
334 *value_function,
335 *exploration_aware,
336 ),
337 DomainStrategy::ScientificComputing {
338 stability_focused,
339 precision_critical,
340 sparse_optimized,
341 } => self.optimize_scientific_computing(
342 context,
343 *stability_focused,
344 *precision_critical,
345 *sparse_optimized,
346 ),
347 }
348 }
349
350 fn optimize_computer_vision(
352 &self,
353 context: &OptimizationContext<A>,
354 resolution_adaptive: bool,
355 batch_norm_tuning: bool,
356 augmentation_aware: bool,
357 ) -> Result<DomainOptimizationConfig<A>> {
358 let mut config = DomainOptimizationConfig::default();
359
360 if resolution_adaptive {
362 let resolution_factor = self.estimate_resolution_factor(&context.problem_chars);
363 config.learning_rate = self.config.base_learning_rate * cast_scalar(resolution_factor)?;
364
365 if context.problem_chars.input_dim > 512 * 512 {
367 config.learning_rate = config.learning_rate * cast_scalar(0.5)?;
368 }
369 }
370
371 if batch_norm_tuning {
373 config.optimizer_type = OptimizerType::AdamW; config
375 .specialized_params
376 .insert("batch_norm_momentum".to_string(), cast_scalar(0.99)?);
377 config
378 .specialized_params
379 .insert("batch_norm_eps".to_string(), cast_scalar(1e-5)?);
380 }
381
382 if augmentation_aware {
384 config.regularization_strength = config.regularization_strength * cast_scalar(1.5)?;
386 config
387 .specialized_params
388 .insert("mixup_alpha".to_string(), cast_scalar(0.2)?);
389 config
390 .specialized_params
391 .insert("cutmix_alpha".to_string(), cast_scalar(1.0)?);
392 }
393
394 config.batch_size = self.select_cv_batch_size(&context.resource_constraints);
396 config.gradient_clip_norm = Some(cast_scalar(1.0)?);
397
398 config.lr_schedule = LearningRateScheduleType::CosineAnnealing {
400 t_max: context.training_config.max_epochs,
401 };
402
403 Ok(config)
404 }
405
406 fn optimize_natural_language(
408 &self,
409 context: &OptimizationContext<A>,
410 sequence_adaptive: bool,
411 attention_optimized: bool,
412 vocab_aware: bool,
413 ) -> Result<DomainOptimizationConfig<A>> {
414 let mut config = DomainOptimizationConfig::default();
415
416 if sequence_adaptive {
418 let seq_length = context.problem_chars.input_dim; if seq_length > 512 {
422 config.learning_rate = self.config.base_learning_rate * cast_scalar(0.7)?;
423 config.gradient_clip_norm = Some(cast_scalar(0.5)?);
424 } else {
425 config.learning_rate = self.config.base_learning_rate;
426 config.gradient_clip_norm = Some(cast_scalar(1.0)?);
427 }
428 }
429
430 if attention_optimized {
432 config.optimizer_type = OptimizerType::AdamW; config
434 .specialized_params
435 .insert("attention_dropout".to_string(), cast_scalar(0.1)?);
436 config
437 .specialized_params
438 .insert("attention_head_dim".to_string(), cast_scalar(64.0)?);
439
440 config
442 .specialized_params
443 .insert("layer_decay_rate".to_string(), cast_scalar(0.95)?);
444 }
445
446 if vocab_aware {
448 let vocab_size = context.problem_chars.output_dim; if vocab_size > 30000 {
452 config
453 .specialized_params
454 .insert("tie_embeddings".to_string(), cast_scalar(1.0)?);
455 config
456 .specialized_params
457 .insert("embedding_dropout".to_string(), cast_scalar(0.1)?);
458 }
459 }
460
461 config.batch_size = self.select_nlp_batch_size(&context.resource_constraints);
463 config.lr_schedule = LearningRateScheduleType::OneCycle {
464 max_lr: config.learning_rate.to_f64().ok_or_else(|| {
465 OptimError::InvalidConfig(
466 "optimize_natural_language: learning_rate must convert to f64".to_string(),
467 )
468 })?,
469 };
470
471 config
473 .specialized_params
474 .insert("warmup_steps".to_string(), cast_scalar(1000.0)?);
475
476 Ok(config)
477 }
478
479 fn optimize_recommendation_systems(
481 &self,
482 context: &OptimizationContext<A>,
483 collaborative_filtering: bool,
484 matrix_factorization: bool,
485 cold_start_aware: bool,
486 ) -> Result<DomainOptimizationConfig<A>> {
487 let mut config = DomainOptimizationConfig::default();
488
489 if collaborative_filtering {
491 config.optimizer_type = OptimizerType::Adam; config.regularization_strength = cast_scalar(0.01)?; config
494 .specialized_params
495 .insert("negative_sampling_rate".to_string(), cast_scalar(5.0)?);
496 }
497
498 if matrix_factorization {
500 config.learning_rate = cast_scalar(0.01)?; config
502 .specialized_params
503 .insert("embedding_dim".to_string(), cast_scalar(128.0)?);
504 config
505 .specialized_params
506 .insert("factorization_rank".to_string(), cast_scalar(50.0)?);
507 }
508
509 if cold_start_aware {
511 config
512 .specialized_params
513 .insert("content_weight".to_string(), cast_scalar(0.3)?);
514 config
515 .specialized_params
516 .insert("popularity_bias".to_string(), cast_scalar(0.1)?);
517 }
518
519 config.batch_size = self.select_recsys_batch_size(&context.resource_constraints);
521 config.gradient_clip_norm = Some(cast_scalar(5.0)?); Ok(config)
524 }
525
526 fn optimize_time_series(
528 &self,
529 context: &OptimizationContext<A>,
530 temporal_aware: bool,
531 seasonality_adaptive: bool,
532 multi_step: bool,
533 ) -> Result<DomainOptimizationConfig<A>> {
534 let mut config = DomainOptimizationConfig::default();
535
536 if temporal_aware {
538 config.optimizer_type = OptimizerType::RMSprop; config.learning_rate = cast_scalar(0.001)?; config.specialized_params.insert(
541 "sequence_length".to_string(),
542 cast_scalar(context.problem_chars.input_dim as f64)?,
543 );
544 }
545
546 if seasonality_adaptive {
548 config
549 .specialized_params
550 .insert("seasonal_periods".to_string(), cast_scalar(24.0)?); config
552 .specialized_params
553 .insert("trend_strength".to_string(), cast_scalar(0.1)?);
554 }
555
556 if multi_step {
558 config
559 .specialized_params
560 .insert("prediction_horizon".to_string(), cast_scalar(12.0)?);
561 config
562 .specialized_params
563 .insert("multi_step_loss_weight".to_string(), cast_scalar(0.8)?);
564 }
565
566 config.batch_size = 32; config.gradient_clip_norm = Some(cast_scalar(1.0)?);
569 config.lr_schedule = LearningRateScheduleType::ReduceOnPlateau {
570 patience: 10,
571 factor: 0.5,
572 };
573
574 Ok(config)
575 }
576
577 fn optimize_reinforcement_learning(
579 &self,
580 context: &OptimizationContext<A>,
581 policy_gradient: bool,
582 value_function: bool,
583 exploration_aware: bool,
584 ) -> Result<DomainOptimizationConfig<A>> {
585 let mut config = DomainOptimizationConfig::default();
586
587 if policy_gradient {
589 config.optimizer_type = OptimizerType::Adam;
590 config.learning_rate = cast_scalar(3e-4)?; config
592 .specialized_params
593 .insert("entropy_coeff".to_string(), cast_scalar(0.01)?);
594 }
595
596 if value_function {
598 config
599 .specialized_params
600 .insert("value_loss_coeff".to_string(), cast_scalar(0.5)?);
601 config
602 .specialized_params
603 .insert("huber_loss_delta".to_string(), cast_scalar(1.0)?);
604 }
605
606 if exploration_aware {
608 config
609 .specialized_params
610 .insert("epsilon_start".to_string(), cast_scalar(1.0)?);
611 config
612 .specialized_params
613 .insert("epsilon_end".to_string(), cast_scalar(0.1)?);
614 config
615 .specialized_params
616 .insert("epsilon_decay".to_string(), cast_scalar(0.995)?);
617 }
618
619 config.batch_size = self.select_rl_batch_size(&context.resource_constraints);
621 config.gradient_clip_norm = Some(cast_scalar(0.5)?); config.lr_schedule = LearningRateScheduleType::Constant; Ok(config)
625 }
626
627 fn optimize_scientific_computing(
629 &self,
630 context: &OptimizationContext<A>,
631 stability_focused: bool,
632 precision_critical: bool,
633 sparse_optimized: bool,
634 ) -> Result<DomainOptimizationConfig<A>> {
635 let mut config = DomainOptimizationConfig::default();
636
637 if stability_focused {
639 config.optimizer_type = OptimizerType::LBFGS; config.learning_rate = cast_scalar(0.1)?; config
642 .specialized_params
643 .insert("line_search_tolerance".to_string(), cast_scalar(1e-6)?);
644 }
645
646 if precision_critical {
648 config
649 .specialized_params
650 .insert("convergence_tolerance".to_string(), cast_scalar(1e-8)?);
651 config
652 .specialized_params
653 .insert("max_iterations".to_string(), cast_scalar(1000.0)?);
654 }
655
656 if sparse_optimized {
658 config.optimizer_type = OptimizerType::Adam;
659 config
660 .specialized_params
661 .insert("sparsity_threshold".to_string(), cast_scalar(1e-6)?);
662 }
663
664 config.batch_size = context.problem_chars.dataset_size.min(1024); config.gradient_clip_norm = None; config.lr_schedule = LearningRateScheduleType::Constant; Ok(config)
670 }
671
672 pub fn update_domain_performance(
674 &mut self,
675 domain: String,
676 metrics: DomainPerformanceMetrics<A>,
677 ) {
678 self.domain_performance
679 .entry(domain)
680 .or_default()
681 .push(metrics);
682 }
683
684 pub fn record_transfer_knowledge(&mut self, knowledge: CrossDomainKnowledge<A>) {
686 self.transfer_knowledge.push(knowledge);
687 }
688
689 pub fn get_domain_recommendations(&self, domain: &str) -> Vec<DomainRecommendation<A>> {
691 let mut recommendations = Vec::new();
692
693 if let Some(history) = self.domain_performance.get(domain) {
695 if !history.is_empty() {
696 let history_len: A = A::from(history.len())
697 .expect("get_domain_recommendations: history.len() must fit in A");
698 let avg_performance =
699 history.iter().map(|m| m.validation_accuracy).sum::<A>() / history_len;
700
701 recommendations.push(DomainRecommendation {
702 recommendation_type: RecommendationType::PerformanceBaseline,
703 description: format!(
704 "Historical average performance: {:.4}",
705 avg_performance.to_f64().expect(
706 "get_domain_recommendations: A (f32/f64) always converts to f64"
707 )
708 ),
709 confidence: A::from(0.8)
710 .expect("get_domain_recommendations: literal 0.8 must fit in A"),
711 action: "Consider this as baseline for improvements".to_string(),
712 });
713 }
714 }
715
716 for knowledge in &self.transfer_knowledge {
718 if knowledge.target_domain == domain {
719 recommendations.push(DomainRecommendation {
720 recommendation_type: RecommendationType::TransferLearning,
721 description: format!(
722 "Transfer from {} domain with {:.2} effectiveness",
723 knowledge.source_domain,
724 knowledge.transfer_score.to_f64().expect(
725 "get_domain_recommendations: A (f32/f64) always converts to f64"
726 )
727 ),
728 confidence: knowledge.transfer_score,
729 action: format!("Use {:?} optimizer", knowledge.successful_strategy),
730 });
731 }
732 }
733
734 recommendations
735 }
736
737 fn estimate_resolution_factor(&self, problem_chars: &ProblemCharacteristics) -> f64 {
739 let resolution = problem_chars.input_dim as f64;
740
741 if resolution > 1_000_000.0 {
742 0.5
744 } else if resolution > 250_000.0 {
745 0.7
747 } else if resolution > 50_000.0 {
748 0.9
750 } else {
751 1.0
753 }
754 }
755
756 fn select_cv_batch_size(&self, constraints: &ResourceConstraints<A>) -> usize {
757 if constraints.max_memory > 16_000_000_000 {
758 128
760 } else if constraints.max_memory > 8_000_000_000 {
761 64
763 } else {
764 32
765 }
766 }
767
768 fn select_nlp_batch_size(&self, constraints: &ResourceConstraints<A>) -> usize {
769 if constraints.max_memory > 32_000_000_000 {
770 64
772 } else if constraints.max_memory > 16_000_000_000 {
773 32
775 } else {
776 16
777 }
778 }
779
780 fn select_recsys_batch_size(&self, constraints: &ResourceConstraints<A>) -> usize {
781 if constraints.max_memory > 8_000_000_000 {
783 512
784 } else {
785 256
786 }
787 }
788
789 fn select_rl_batch_size(&self, constraints: &ResourceConstraints<A>) -> usize {
790 if constraints.max_memory > 16_000_000_000 {
794 128
796 } else if constraints.max_memory > 8_000_000_000 {
797 64
799 } else {
800 32
801 }
802 }
803
804 fn default_config_for_strategy(strategy: &DomainStrategy) -> DomainConfig<A> {
806 match strategy {
807 DomainStrategy::ComputerVision { .. } => DomainConfig {
808 base_learning_rate: A::from(0.001)
809 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
810 recommended_batch_sizes: vec![32, 64, 128],
811 gradient_clip_values: vec![
812 A::from(1.0)
813 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
814 A::from(2.0)
815 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
816 ],
817 regularization_range: (
818 A::from(1e-5)
819 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
820 A::from(1e-2)
821 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
822 ),
823 optimizer_ranking: vec![
824 OptimizerType::AdamW,
825 OptimizerType::SGDMomentum,
826 OptimizerType::Adam,
827 ],
828 domain_params: HashMap::new(),
829 },
830 DomainStrategy::NaturalLanguage { .. } => DomainConfig {
831 base_learning_rate: A::from(2e-5)
832 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
833 recommended_batch_sizes: vec![16, 32, 64],
834 gradient_clip_values: vec![
835 A::from(0.5)
836 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
837 A::from(1.0)
838 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
839 ],
840 regularization_range: (
841 A::from(1e-4)
842 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
843 A::from(1e-1)
844 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
845 ),
846 optimizer_ranking: vec![OptimizerType::AdamW, OptimizerType::Adam],
847 domain_params: HashMap::new(),
848 },
849 DomainStrategy::RecommendationSystems { .. } => DomainConfig {
850 base_learning_rate: A::from(0.01)
851 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
852 recommended_batch_sizes: vec![128, 256, 512],
853 gradient_clip_values: vec![
854 A::from(5.0)
855 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
856 A::from(10.0)
857 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
858 ],
859 regularization_range: (
860 A::from(1e-3)
861 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
862 A::from(1e-1)
863 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
864 ),
865 optimizer_ranking: vec![OptimizerType::Adam, OptimizerType::AdaGrad],
866 domain_params: HashMap::new(),
867 },
868 DomainStrategy::TimeSeries { .. } => DomainConfig {
869 base_learning_rate: A::from(0.001)
870 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
871 recommended_batch_sizes: vec![16, 32, 64],
872 gradient_clip_values: vec![A::from(1.0)
873 .expect("default_config_for_strategy: literal must fit in A (f32/f64)")],
874 regularization_range: (
875 A::from(1e-4)
876 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
877 A::from(1e-2)
878 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
879 ),
880 optimizer_ranking: vec![OptimizerType::RMSprop, OptimizerType::Adam],
881 domain_params: HashMap::new(),
882 },
883 DomainStrategy::ReinforcementLearning { .. } => DomainConfig {
884 base_learning_rate: A::from(3e-4)
885 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
886 recommended_batch_sizes: vec![32, 64, 128],
887 gradient_clip_values: vec![A::from(0.5)
888 .expect("default_config_for_strategy: literal must fit in A (f32/f64)")],
889 regularization_range: (
890 A::from(1e-4)
891 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
892 A::from(1e-2)
893 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
894 ),
895 optimizer_ranking: vec![OptimizerType::Adam],
896 domain_params: HashMap::new(),
897 },
898 DomainStrategy::ScientificComputing { .. } => DomainConfig {
899 base_learning_rate: A::from(0.1)
900 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
901 recommended_batch_sizes: vec![64, 128, 256, 512],
902 gradient_clip_values: vec![],
903 regularization_range: (
904 A::from(1e-6)
905 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
906 A::from(1e-3)
907 .expect("default_config_for_strategy: literal must fit in A (f32/f64)"),
908 ),
909 optimizer_ranking: vec![OptimizerType::LBFGS, OptimizerType::Adam],
910 domain_params: HashMap::new(),
911 },
912 }
913 }
914}
915
916#[derive(Debug, Clone)]
918pub struct DomainOptimizationConfig<A: Float> {
919 pub optimizer_type: OptimizerType,
921 pub learning_rate: A,
923 pub batch_size: usize,
925 pub gradient_clip_norm: Option<A>,
927 pub regularization_strength: A,
929 pub lr_schedule: LearningRateScheduleType,
931 pub specialized_params: HashMap<String, A>,
933}
934
935impl<A: Float + Send + Sync> Default for DomainOptimizationConfig<A> {
936 fn default() -> Self {
937 Self {
938 optimizer_type: OptimizerType::Adam,
939 learning_rate: A::from(0.001)
940 .expect("DomainOptimizationConfig: default learning_rate (0.001) must fit in A"),
941 batch_size: 32,
942 gradient_clip_norm: Some(A::from(1.0).expect(
943 "DomainOptimizationConfig: default gradient_clip_norm (1.0) must fit in A",
944 )),
945 regularization_strength: A::from(1e-4).expect(
946 "DomainOptimizationConfig: default regularization_strength (1e-4) must fit in A",
947 ),
948 lr_schedule: LearningRateScheduleType::Constant,
949 specialized_params: HashMap::new(),
950 }
951 }
952}
953
954#[derive(Debug, Clone)]
956pub struct DomainRecommendation<A: Float> {
957 pub recommendation_type: RecommendationType,
959 pub description: String,
961 pub confidence: A,
963 pub action: String,
965}
966
967#[derive(Debug, Clone)]
969pub enum RecommendationType {
970 PerformanceBaseline,
972 TransferLearning,
974 HyperparameterTuning,
976 ArchitectureChange,
978 ResourceOptimization,
980}
981
982#[cfg(test)]
983mod tests {
984 use super::*;
985 use crate::adaptive_selection::ProblemType;
986
987 #[test]
988 fn test_domain_specific_selector_creation() {
989 let strategy = DomainStrategy::ComputerVision {
990 resolution_adaptive: true,
991 batch_norm_tuning: true,
992 augmentation_aware: true,
993 };
994
995 let selector = DomainSpecificSelector::<f64>::new(strategy);
996 assert_eq!(selector.config.optimizer_ranking[0], OptimizerType::AdamW);
997 }
998
999 #[test]
1000 fn test_computer_vision_optimization() {
1001 let strategy = DomainStrategy::ComputerVision {
1002 resolution_adaptive: true,
1003 batch_norm_tuning: true,
1004 augmentation_aware: true,
1005 };
1006
1007 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1008
1009 let context = OptimizationContext {
1010 problem_chars: ProblemCharacteristics {
1011 dataset_size: 50000,
1012 input_dim: 224 * 224 * 3, output_dim: 1000,
1014 problem_type: ProblemType::ComputerVision,
1015 gradient_sparsity: 0.1,
1016 gradient_noise: 0.05,
1017 memory_budget: 8_000_000_000,
1018 time_budget: 3600.0,
1019 batch_size: 64,
1020 lr_sensitivity: 0.5,
1021 regularization_strength: 0.01,
1022 architecture_type: Some("ResNet".to_string()),
1023 },
1024 resource_constraints: ResourceConstraints {
1025 max_memory: 17_000_000_000, max_time: 7200.0,
1027 gpu_available: true,
1028 distributed_capable: false,
1029 energy_efficient: false,
1030 },
1031 training_config: TrainingConfiguration {
1032 max_epochs: 100,
1033 early_stopping_patience: 10,
1034 validation_frequency: 1,
1035 lr_schedule_type: LearningRateScheduleType::CosineAnnealing { t_max: 100 },
1036 regularization_approach: RegularizationApproach::L2Only { weight: 1e-4 },
1037 },
1038 domain_metadata: HashMap::new(),
1039 };
1040
1041 selector.setcontext(context);
1042 let config = selector
1043 .select_optimal_config()
1044 .expect("selector.select_optimal_config succeeds in test_computer_vision_optimization");
1045
1046 assert_eq!(config.optimizer_type, OptimizerType::AdamW);
1047 assert_eq!(config.batch_size, 128); assert!(config.gradient_clip_norm.is_some());
1049 }
1050
1051 #[test]
1052 fn test_reinforcement_learning_optimization_uses_resource_constraints() {
1053 let strategy = DomainStrategy::ReinforcementLearning {
1054 policy_gradient: true,
1055 value_function: true,
1056 exploration_aware: true,
1057 };
1058
1059 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1060
1061 let make_context = |max_memory: usize| OptimizationContext {
1062 problem_chars: ProblemCharacteristics {
1063 dataset_size: 10000,
1064 input_dim: 128,
1065 output_dim: 4,
1066 problem_type: ProblemType::ReinforcementLearning,
1067 gradient_sparsity: 0.0,
1068 gradient_noise: 0.1,
1069 memory_budget: max_memory,
1070 time_budget: 3600.0,
1071 batch_size: 64,
1072 lr_sensitivity: 0.5,
1073 regularization_strength: 0.0,
1074 architecture_type: None,
1075 },
1076 resource_constraints: ResourceConstraints {
1077 max_memory,
1078 max_time: 3600.0,
1079 gpu_available: true,
1080 distributed_capable: false,
1081 energy_efficient: false,
1082 },
1083 training_config: TrainingConfiguration {
1084 max_epochs: 1000,
1085 early_stopping_patience: 50,
1086 validation_frequency: 10,
1087 lr_schedule_type: LearningRateScheduleType::Constant,
1088 regularization_approach: RegularizationApproach::L2Only { weight: 0.0 },
1089 },
1090 domain_metadata: HashMap::new(),
1091 };
1092
1093 selector.setcontext(make_context(4_000_000_000));
1095 let low_mem_config = selector.select_optimal_config().expect("selector.select_optimal_config succeeds in test_reinforcement_learning_optimization_uses_resource_constraints");
1096 assert_eq!(low_mem_config.batch_size, 32);
1097
1098 selector.setcontext(make_context(17_000_000_000));
1102 let high_mem_config = selector.select_optimal_config().expect("selector.select_optimal_config succeeds in test_reinforcement_learning_optimization_uses_resource_constraints");
1103 assert_eq!(high_mem_config.batch_size, 128);
1104 assert!(high_mem_config.gradient_clip_norm.is_some());
1105 }
1106
1107 #[test]
1108 fn test_natural_language_optimization() {
1109 let strategy = DomainStrategy::NaturalLanguage {
1110 sequence_adaptive: true,
1111 attention_optimized: true,
1112 vocab_aware: true,
1113 };
1114
1115 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1116
1117 let context = OptimizationContext {
1118 problem_chars: ProblemCharacteristics {
1119 dataset_size: 100000,
1120 input_dim: 512, output_dim: 50000, problem_type: ProblemType::NaturalLanguage,
1123 gradient_sparsity: 0.2,
1124 gradient_noise: 0.1,
1125 memory_budget: 32_000_000_000,
1126 time_budget: 7200.0,
1127 batch_size: 32,
1128 lr_sensitivity: 0.8,
1129 regularization_strength: 0.1,
1130 architecture_type: Some("Transformer".to_string()),
1131 },
1132 resource_constraints: ResourceConstraints {
1133 max_memory: 32_000_000_000,
1134 max_time: 10800.0,
1135 gpu_available: true,
1136 distributed_capable: true,
1137 energy_efficient: false,
1138 },
1139 training_config: TrainingConfiguration {
1140 max_epochs: 50,
1141 early_stopping_patience: 5,
1142 validation_frequency: 1,
1143 lr_schedule_type: LearningRateScheduleType::OneCycle { max_lr: 2e-5 },
1144 regularization_approach: RegularizationApproach::Dropout { dropout_rate: 0.1 },
1145 },
1146 domain_metadata: HashMap::new(),
1147 };
1148
1149 selector.setcontext(context);
1150 let config = selector.select_optimal_config().expect(
1151 "selector.select_optimal_config succeeds in test_natural_language_optimization",
1152 );
1153
1154 assert_eq!(config.optimizer_type, OptimizerType::AdamW);
1155 assert!(config.specialized_params.contains_key("warmup_steps"));
1156 assert!(config.specialized_params.contains_key("tie_embeddings")); }
1158
1159 #[test]
1160 fn test_time_series_optimization() {
1161 let strategy = DomainStrategy::TimeSeries {
1162 temporal_aware: true,
1163 seasonality_adaptive: true,
1164 multi_step: true,
1165 };
1166
1167 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1168
1169 let context = OptimizationContext {
1170 problem_chars: ProblemCharacteristics {
1171 dataset_size: 10000,
1172 input_dim: 168, output_dim: 24, problem_type: ProblemType::TimeSeries,
1175 gradient_sparsity: 0.05,
1176 gradient_noise: 0.2,
1177 memory_budget: 4_000_000_000,
1178 time_budget: 1800.0,
1179 batch_size: 32,
1180 lr_sensitivity: 0.7,
1181 regularization_strength: 0.01,
1182 architecture_type: Some("LSTM".to_string()),
1183 },
1184 resource_constraints: ResourceConstraints {
1185 max_memory: 8_000_000_000,
1186 max_time: 3600.0,
1187 gpu_available: true,
1188 distributed_capable: false,
1189 energy_efficient: true,
1190 },
1191 training_config: TrainingConfiguration {
1192 max_epochs: 200,
1193 early_stopping_patience: 20,
1194 validation_frequency: 5,
1195 lr_schedule_type: LearningRateScheduleType::ReduceOnPlateau {
1196 patience: 10,
1197 factor: 0.5,
1198 },
1199 regularization_approach: RegularizationApproach::L2Only { weight: 1e-4 },
1200 },
1201 domain_metadata: HashMap::new(),
1202 };
1203
1204 selector.setcontext(context);
1205 let config = selector
1206 .select_optimal_config()
1207 .expect("selector.select_optimal_config succeeds in test_time_series_optimization");
1208
1209 assert_eq!(config.optimizer_type, OptimizerType::RMSprop);
1210 assert_eq!(config.batch_size, 32);
1211 assert!(config.specialized_params.contains_key("seasonal_periods"));
1212 assert!(config.specialized_params.contains_key("prediction_horizon"));
1213 }
1214
1215 #[test]
1216 fn test_performance_tracking() {
1217 let strategy = DomainStrategy::ComputerVision {
1218 resolution_adaptive: true,
1219 batch_norm_tuning: false,
1220 augmentation_aware: false,
1221 };
1222
1223 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1224
1225 let metrics = DomainPerformanceMetrics {
1226 validation_accuracy: 0.95,
1227 domain_specific_score: 0.92,
1228 stability_score: 0.88,
1229 convergence_epochs: 50,
1230 resource_efficiency: 0.85,
1231 transfer_score: 0.7,
1232 };
1233
1234 selector.update_domain_performance("computer_vision".to_string(), metrics);
1235
1236 let recommendations = selector.get_domain_recommendations("computer_vision");
1237 assert!(!recommendations.is_empty());
1238 assert!(recommendations[0].description.contains("0.95"));
1239 }
1240
1241 #[test]
1242 fn test_cross_domain_transfer() {
1243 let strategy = DomainStrategy::ComputerVision {
1244 resolution_adaptive: true,
1245 batch_norm_tuning: true,
1246 augmentation_aware: true,
1247 };
1248
1249 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1250
1251 let transfer_knowledge = CrossDomainKnowledge {
1252 source_domain: "natural_language".to_string(),
1253 target_domain: "computer_vision".to_string(),
1254 transferable_params: HashMap::from([
1255 ("learning_rate".to_string(), 0.001),
1256 ("weight_decay".to_string(), 0.01),
1257 ]),
1258 transfer_score: 0.8,
1259 successful_strategy: OptimizerType::AdamW,
1260 };
1261
1262 selector.record_transfer_knowledge(transfer_knowledge);
1263
1264 let recommendations = selector.get_domain_recommendations("computer_vision");
1265 assert!(recommendations
1266 .iter()
1267 .any(|r| matches!(r.recommendation_type, RecommendationType::TransferLearning)));
1268 }
1269
1270 #[test]
1271 fn test_scientific_computing_optimization() {
1272 let strategy = DomainStrategy::ScientificComputing {
1273 stability_focused: true,
1274 precision_critical: true,
1275 sparse_optimized: false,
1276 };
1277
1278 let mut selector = DomainSpecificSelector::<f64>::new(strategy);
1279
1280 let context = OptimizationContext {
1281 problem_chars: ProblemCharacteristics {
1282 dataset_size: 1000,
1283 input_dim: 100,
1284 output_dim: 1,
1285 problem_type: ProblemType::Regression,
1286 gradient_sparsity: 0.01,
1287 gradient_noise: 0.01,
1288 memory_budget: 16_000_000_000,
1289 time_budget: 7200.0,
1290 batch_size: 100,
1291 lr_sensitivity: 0.3,
1292 regularization_strength: 1e-6,
1293 architecture_type: Some("MLP".to_string()),
1294 },
1295 resource_constraints: ResourceConstraints {
1296 max_memory: 16_000_000_000,
1297 max_time: 7200.0,
1298 gpu_available: false,
1299 distributed_capable: false,
1300 energy_efficient: false,
1301 },
1302 training_config: TrainingConfiguration {
1303 max_epochs: 1000,
1304 early_stopping_patience: 100,
1305 validation_frequency: 10,
1306 lr_schedule_type: LearningRateScheduleType::Constant,
1307 regularization_approach: RegularizationApproach::L2Only { weight: 1e-6 },
1308 },
1309 domain_metadata: HashMap::new(),
1310 };
1311
1312 selector.setcontext(context);
1313 let config = selector.select_optimal_config().expect(
1314 "selector.select_optimal_config succeeds in test_scientific_computing_optimization",
1315 );
1316
1317 assert_eq!(config.optimizer_type, OptimizerType::LBFGS);
1318 assert!(config.gradient_clip_norm.is_none()); assert!(config
1320 .specialized_params
1321 .contains_key("convergence_tolerance"));
1322 }
1323}