Skip to main content

optirs_core/domain_specific/
mod.rs

1// Domain-specific optimization strategies
2//
3// This module provides specialized optimization strategies tailored for different
4// machine learning domains, building on the adaptive selection framework to provide
5// domain-aware optimization approaches.
6
7use 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
14/// Fallibly converts an `f64`/`usize` literal or measurement into the
15/// selector's generic scalar type `A`.
16///
17/// Centralizes what used to be `A::from(x).expect("unwrap failed")` call
18/// sites across the domain-specific hyperparameter builders: instead of
19/// panicking, a type that genuinely cannot represent `x` now produces an
20/// honest [`OptimError`].
21fn 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/// Domain-specific optimization strategy
30#[derive(Debug, Clone)]
31pub enum DomainStrategy {
32    /// Computer Vision optimization
33    ComputerVision {
34        /// Image resolution considerations
35        resolution_adaptive: bool,
36        /// Batch normalization optimization
37        batch_norm_tuning: bool,
38        /// Data augmentation awareness
39        augmentation_aware: bool,
40    },
41    /// Natural Language Processing optimization
42    NaturalLanguage {
43        /// Sequence length adaptation
44        sequence_adaptive: bool,
45        /// Attention mechanism optimization
46        attention_optimized: bool,
47        /// Vocabulary size considerations
48        vocab_aware: bool,
49    },
50    /// Recommendation Systems optimization
51    RecommendationSystems {
52        /// Collaborative filtering optimization
53        collaborative_filtering: bool,
54        /// Matrix factorization tuning
55        matrix_factorization: bool,
56        /// Cold start handling
57        cold_start_aware: bool,
58    },
59    /// Time Series optimization
60    TimeSeries {
61        /// Temporal dependency handling
62        temporal_aware: bool,
63        /// Seasonality consideration
64        seasonality_adaptive: bool,
65        /// Multi-step ahead optimization
66        multi_step: bool,
67    },
68    /// Reinforcement Learning optimization
69    ReinforcementLearning {
70        /// Policy gradient optimization
71        policy_gradient: bool,
72        /// Value function optimization
73        value_function: bool,
74        /// Exploration-exploitation balance
75        exploration_aware: bool,
76    },
77    /// Scientific Computing optimization
78    ScientificComputing {
79        /// Numerical stability prioritization
80        stability_focused: bool,
81        /// High precision requirements
82        precision_critical: bool,
83        /// Sparse matrix optimization
84        sparse_optimized: bool,
85    },
86}
87
88/// Domain-specific configuration parameters
89#[derive(Debug, Clone)]
90pub struct DomainConfig<A: Float> {
91    /// Base learning rate for the domain
92    pub base_learning_rate: A,
93    /// Batch size recommendations
94    pub recommended_batch_sizes: Vec<usize>,
95    /// Gradient clipping thresholds
96    pub gradient_clip_values: Vec<A>,
97    /// Regularization strengths
98    pub regularization_range: (A, A),
99    /// Optimizer preferences (ranked by effectiveness)
100    pub optimizer_ranking: Vec<OptimizerType>,
101    /// Domain-specific hyperparameters
102    pub domain_params: HashMap<String, A>,
103}
104
105/// Domain-specific optimizer selector
106#[derive(Debug)]
107pub struct DomainSpecificSelector<A: Float> {
108    /// Current domain strategy
109    strategy: DomainStrategy,
110    /// Domain configuration
111    config: DomainConfig<A>,
112    /// Performance history per domain
113    domain_performance: HashMap<String, Vec<DomainPerformanceMetrics<A>>>,
114    /// Cross-domain transfer learning data
115    transfer_knowledge: Vec<CrossDomainKnowledge<A>>,
116    /// Current optimization context
117    currentcontext: Option<OptimizationContext<A>>,
118}
119
120/// Performance metrics specific to domains
121#[derive(Debug, Clone)]
122pub struct DomainPerformanceMetrics<A: Float> {
123    /// Standard performance metrics
124    pub validation_accuracy: A,
125    /// Domain-specific metrics
126    pub domain_specific_score: A,
127    /// Training stability
128    pub stability_score: A,
129    /// Convergence speed (epochs to target)
130    pub convergence_epochs: usize,
131    /// Resource efficiency
132    pub resource_efficiency: A,
133    /// Transfer learning potential
134    pub transfer_score: A,
135}
136
137/// Cross-domain knowledge transfer
138#[derive(Debug, Clone)]
139pub struct CrossDomainKnowledge<A: Float> {
140    /// Source domain
141    pub source_domain: String,
142    /// Target domain
143    pub target_domain: String,
144    /// Transferable hyperparameters
145    pub transferable_params: HashMap<String, A>,
146    /// Transfer effectiveness score
147    pub transfer_score: A,
148    /// Optimization strategy that worked
149    pub successful_strategy: OptimizerType,
150}
151
152/// Current optimization context
153#[derive(Debug, Clone)]
154pub struct OptimizationContext<A: Float> {
155    /// Problem characteristics
156    pub problem_chars: ProblemCharacteristics,
157    /// Resource constraints
158    pub resource_constraints: ResourceConstraints<A>,
159    /// Training characteristics
160    pub training_config: TrainingConfiguration<A>,
161    /// Domain-specific metadata
162    pub domain_metadata: HashMap<String, String>,
163}
164
165/// Resource constraints for optimization
166#[derive(Debug, Clone)]
167pub struct ResourceConstraints<A: Float> {
168    /// Maximum memory available (bytes)
169    pub max_memory: usize,
170    /// Maximum training time (seconds)
171    pub max_time: A,
172    /// GPU availability and type
173    pub gpu_available: bool,
174    /// Distributed training capability
175    pub distributed_capable: bool,
176    /// Energy efficiency requirements
177    pub energy_efficient: bool,
178}
179
180/// Training configuration parameters
181#[derive(Debug, Clone)]
182pub struct TrainingConfiguration<A: Float> {
183    /// Maximum number of epochs
184    pub max_epochs: usize,
185    /// Early stopping patience
186    pub early_stopping_patience: usize,
187    /// Validation frequency
188    pub validation_frequency: usize,
189    /// Learning rate scheduling
190    pub lr_schedule_type: LearningRateScheduleType,
191    /// Regularization approach
192    pub regularization_approach: RegularizationApproach<A>,
193}
194
195/// Learning rate schedule types
196#[derive(Debug, Clone)]
197pub enum LearningRateScheduleType {
198    /// Constant learning rate
199    Constant,
200    /// Exponential decay
201    ExponentialDecay {
202        /// Decay rate
203        decay_rate: f64,
204    },
205    /// Cosine annealing
206    CosineAnnealing {
207        /// Maximum number of iterations
208        t_max: usize,
209    },
210    /// Reduce on plateau
211    ReduceOnPlateau {
212        /// Number of epochs with no improvement
213        patience: usize,
214        /// Factor by which learning rate will be reduced
215        factor: f64,
216    },
217    /// One cycle policy
218    OneCycle {
219        /// Maximum learning rate
220        max_lr: f64,
221    },
222}
223
224/// Regularization approach
225#[derive(Debug, Clone)]
226pub enum RegularizationApproach<A: Float> {
227    /// L2 regularization only
228    L2Only {
229        /// Regularization weight
230        weight: A,
231    },
232    /// L1 regularization only
233    L1Only {
234        /// Regularization weight
235        weight: A,
236    },
237    /// Elastic net (L1 + L2)
238    ElasticNet {
239        /// L1 regularization weight
240        l1_weight: A,
241        /// L2 regularization weight
242        l2_weight: A,
243    },
244    /// Dropout regularization
245    Dropout {
246        /// Dropout rate
247        dropout_rate: A,
248    },
249    /// Combined approach
250    Combined {
251        /// L2 regularization weight
252        l2_weight: A,
253        /// Dropout rate
254        dropout_rate: A,
255        /// Additional regularization techniques
256        additional_techniques: Vec<String>,
257    },
258}
259
260impl<A: Float + ScalarOperand + Debug + std::iter::Sum + Send + Sync> DomainSpecificSelector<A> {
261    /// Create a new domain-specific selector
262    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    /// Set optimization context
275    pub fn setcontext(&mut self, context: OptimizationContext<A>) {
276        self.currentcontext = Some(context);
277    }
278
279    /// Select optimal configuration for the current domain and context
280    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    /// Optimize for computer vision tasks
351    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        // Resolution-_adaptive optimization
361        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            // Larger images need smaller learning rates
366            if context.problem_chars.input_dim > 512 * 512 {
367                config.learning_rate = config.learning_rate * cast_scalar(0.5)?;
368            }
369        }
370
371        // Batch normalization _tuning
372        if batch_norm_tuning {
373            config.optimizer_type = OptimizerType::AdamW; // Better for batch norm
374            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        // Data augmentation awareness
383        if augmentation_aware {
384            // More aggressive regularization with augmentation
385            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        // CV-specific optimizations
395        config.batch_size = self.select_cv_batch_size(&context.resource_constraints);
396        config.gradient_clip_norm = Some(cast_scalar(1.0)?);
397
398        // Use cosine annealing for CV tasks
399        config.lr_schedule = LearningRateScheduleType::CosineAnnealing {
400            t_max: context.training_config.max_epochs,
401        };
402
403        Ok(config)
404    }
405
406    /// Optimize for natural language processing tasks
407    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        // Sequence-_adaptive optimization
417        if sequence_adaptive {
418            let seq_length = context.problem_chars.input_dim; // Assuming input_dim represents sequence length
419
420            // Longer sequences need more careful optimization
421            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        // Attention mechanism optimization
431        if attention_optimized {
432            config.optimizer_type = OptimizerType::AdamW; // Best for transformers
433            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            // Layer-wise learning rate decay for transformers
441            config
442                .specialized_params
443                .insert("layer_decay_rate".to_string(), cast_scalar(0.95)?);
444        }
445
446        // Vocabulary-_aware optimization
447        if vocab_aware {
448            let vocab_size = context.problem_chars.output_dim; // Assuming output_dim represents vocab size
449
450            // Large vocabularies need special handling
451            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        // NLP-specific optimizations
462        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        // Warmup for transformers
472        config
473            .specialized_params
474            .insert("warmup_steps".to_string(), cast_scalar(1000.0)?);
475
476        Ok(config)
477    }
478
479    /// Optimize for recommendation systems
480    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        // Collaborative _filtering optimization
490        if collaborative_filtering {
491            config.optimizer_type = OptimizerType::Adam; // Good for sparse data
492            config.regularization_strength = cast_scalar(0.01)?; // Prevent overfitting
493            config
494                .specialized_params
495                .insert("negative_sampling_rate".to_string(), cast_scalar(5.0)?);
496        }
497
498        // Matrix _factorization tuning
499        if matrix_factorization {
500            config.learning_rate = cast_scalar(0.01)?; // Lower LR for stability
501            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        // Cold start handling
510        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        // RecSys-specific optimizations
520        config.batch_size = self.select_recsys_batch_size(&context.resource_constraints);
521        config.gradient_clip_norm = Some(cast_scalar(5.0)?); // Higher clip for sparse gradients
522
523        Ok(config)
524    }
525
526    /// Optimize for time series tasks
527    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        // Temporal dependency handling
537        if temporal_aware {
538            config.optimizer_type = OptimizerType::RMSprop; // Good for RNNs
539            config.learning_rate = cast_scalar(0.001)?; // Conservative for temporal stability
540            config.specialized_params.insert(
541                "sequence_length".to_string(),
542                cast_scalar(context.problem_chars.input_dim as f64)?,
543            );
544        }
545
546        // Seasonality consideration
547        if seasonality_adaptive {
548            config
549                .specialized_params
550                .insert("seasonal_periods".to_string(), cast_scalar(24.0)?); // Daily pattern
551            config
552                .specialized_params
553                .insert("trend_strength".to_string(), cast_scalar(0.1)?);
554        }
555
556        // Multi-_step ahead optimization
557        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        // Time series-specific optimizations
567        config.batch_size = 32; // Smaller batches for temporal consistency
568        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    /// Optimize for reinforcement learning tasks
578    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        // Policy _gradient optimization
588        if policy_gradient {
589            config.optimizer_type = OptimizerType::Adam;
590            config.learning_rate = cast_scalar(3e-4)?; // Standard RL learning rate
591            config
592                .specialized_params
593                .insert("entropy_coeff".to_string(), cast_scalar(0.01)?);
594        }
595
596        // Value _function optimization
597        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        // Exploration-exploitation balance
607        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        // RL-specific optimizations
620        config.batch_size = self.select_rl_batch_size(&context.resource_constraints);
621        config.gradient_clip_norm = Some(cast_scalar(0.5)?); // Important for RL stability
622        config.lr_schedule = LearningRateScheduleType::Constant; // Often constant in RL
623
624        Ok(config)
625    }
626
627    /// Optimize for scientific computing tasks
628    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        // Numerical stability prioritization
638        if stability_focused {
639            config.optimizer_type = OptimizerType::LBFGS; // More stable for scientific problems
640            config.learning_rate = cast_scalar(0.1)?; // Higher LR for LBFGS
641            config
642                .specialized_params
643                .insert("line_search_tolerance".to_string(), cast_scalar(1e-6)?);
644        }
645
646        // High precision requirements
647        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        // Sparse matrix optimization
657        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        // Scientific computing-specific optimizations
665        config.batch_size = context.problem_chars.dataset_size.min(1024); // Can use larger batches
666        config.gradient_clip_norm = None; // Don't clip for scientific precision
667        config.lr_schedule = LearningRateScheduleType::Constant; // Consistent optimization
668
669        Ok(config)
670    }
671
672    /// Update performance based on training results
673    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    /// Record cross-domain transfer knowledge
685    pub fn record_transfer_knowledge(&mut self, knowledge: CrossDomainKnowledge<A>) {
686        self.transfer_knowledge.push(knowledge);
687    }
688
689    /// Get domain-specific recommendations
690    pub fn get_domain_recommendations(&self, domain: &str) -> Vec<DomainRecommendation<A>> {
691        let mut recommendations = Vec::new();
692
693        // Analyze historical performance for this domain
694        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        // Cross-domain transfer recommendations
717        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    /// Helper methods for domain-specific optimizations
738    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            // Very high resolution
743            0.5
744        } else if resolution > 250_000.0 {
745            // High resolution
746            0.7
747        } else if resolution > 50_000.0 {
748            // Medium resolution
749            0.9
750        } else {
751            // Low resolution
752            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            // 16GB+
759            128
760        } else if constraints.max_memory > 8_000_000_000 {
761            // 8GB+
762            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            // 32GB+
771            64
772        } else if constraints.max_memory > 16_000_000_000 {
773            // 16GB+
774            32
775        } else {
776            16
777        }
778    }
779
780    fn select_recsys_batch_size(&self, constraints: &ResourceConstraints<A>) -> usize {
781        // RecSys can typically use larger batches due to simpler models
782        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        // RL rollout/replay batches scale with available memory, but more
791        // conservatively than supervised CV/NLP batches: most on-policy and
792        // off-policy algorithms are sensitive to very large batch sizes.
793        if constraints.max_memory > 16_000_000_000 {
794            // 16GB+
795            128
796        } else if constraints.max_memory > 8_000_000_000 {
797            // 8GB+
798            64
799        } else {
800            32
801        }
802    }
803
804    /// Create default configuration for a strategy
805    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/// Final domain optimization configuration
917#[derive(Debug, Clone)]
918pub struct DomainOptimizationConfig<A: Float> {
919    /// Selected optimizer type
920    pub optimizer_type: OptimizerType,
921    /// Optimized learning rate
922    pub learning_rate: A,
923    /// Optimized batch size
924    pub batch_size: usize,
925    /// Gradient clipping norm (if applicable)
926    pub gradient_clip_norm: Option<A>,
927    /// Regularization strength
928    pub regularization_strength: A,
929    /// Learning rate schedule
930    pub lr_schedule: LearningRateScheduleType,
931    /// Domain-specific specialized parameters
932    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/// Domain-specific recommendation
955#[derive(Debug, Clone)]
956pub struct DomainRecommendation<A: Float> {
957    /// Type of recommendation
958    pub recommendation_type: RecommendationType,
959    /// Human-readable description
960    pub description: String,
961    /// Confidence in the recommendation (0.0-1.0)
962    pub confidence: A,
963    /// Suggested action
964    pub action: String,
965}
966
967/// Types of domain recommendations
968#[derive(Debug, Clone)]
969pub enum RecommendationType {
970    /// Performance baseline information
971    PerformanceBaseline,
972    /// Transfer learning suggestion
973    TransferLearning,
974    /// Hyperparameter adjustment
975    HyperparameterTuning,
976    /// Architecture modification
977    ArchitectureChange,
978    /// Resource optimization
979    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, // Standard ImageNet resolution
1013                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, // Slightly above 16GB to trigger 128 batch size
1026                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); // Should select larger batch size for high memory
1048        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        // Low memory: falls back to the smallest RL batch-size tier.
1094        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        // High memory: batch size should scale up accordingly, proving the
1099        // resource constraints actually drive the RL batch-size decision
1100        // rather than a hardcoded constant.
1101        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,    // Sequence length
1121                output_dim: 50000, // Large vocabulary
1122                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")); // Large vocab
1157    }
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, // One week of hourly data
1173                output_dim: 24, // Next 24 hours
1174                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()); // No clipping for precision
1319        assert!(config
1320            .specialized_params
1321            .contains_key("convergence_tolerance"));
1322    }
1323}