1use anyhow::Result;
7use chrono::{DateTime, Utc};
8use scirs2_core::ndarray::*; use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11
12#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct AdvancedMLDebuggingConfig {
15 pub enable_layer_wise_lr_analysis: bool,
17 pub enable_model_sensitivity_analysis: bool,
19 pub enable_gradient_flow_optimization: bool,
21 pub enable_neural_architecture_debugging: bool,
23 pub enable_activation_pattern_analysis: bool,
25 pub enable_weight_distribution_analysis: bool,
27 pub enable_training_dynamics_analysis: bool,
29 pub enable_optimization_landscape_analysis: bool,
31 pub sensitivity_samples: usize,
33 pub lr_adaptation_threshold: f64,
35 pub max_layers_to_analyze: usize,
37}
38
39impl Default for AdvancedMLDebuggingConfig {
40 fn default() -> Self {
41 Self {
42 enable_layer_wise_lr_analysis: true,
43 enable_model_sensitivity_analysis: true,
44 enable_gradient_flow_optimization: true,
45 enable_neural_architecture_debugging: true,
46 enable_activation_pattern_analysis: true,
47 enable_weight_distribution_analysis: true,
48 enable_training_dynamics_analysis: true,
49 enable_optimization_landscape_analysis: true,
50 sensitivity_samples: 1000,
51 lr_adaptation_threshold: 0.1,
52 max_layers_to_analyze: 50,
53 }
54 }
55}
56
57#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct LayerWiseLRAnalysisResult {
60 pub timestamp: DateTime<Utc>,
62 pub layer_lr_recommendations: HashMap<String, LayerLRRecommendation>,
64 pub global_lr_insights: GlobalLRInsights,
66 pub adaptation_strategy: LRAdaptationStrategy,
68 pub training_phase_recommendations: Vec<TrainingPhaseRecommendation>,
70 pub lr_schedule_predictions: Vec<LRSchedulePrediction>,
72}
73
74#[derive(Debug, Clone, Serialize, Deserialize)]
76pub struct LayerLRRecommendation {
77 pub layer_id: String,
79 pub layer_type: String,
81 pub current_lr: f64,
83 pub recommended_lr: f64,
85 pub confidence: Option<f64>,
93 pub reasoning: String,
95 pub layer_metrics: LayerLRMetrics,
97 pub lr_sensitivity: f64,
99 pub urgency: AdaptationUrgency,
101}
102
103#[derive(Debug, Clone, Serialize, Deserialize)]
105pub struct LayerLRMetrics {
106 pub gradient_magnitude: f64,
108 pub weight_update_magnitude: f64,
110 pub parameter_norm: f64,
112 pub loss_contribution: f64,
114 pub stability_score: f64,
116 pub convergence_rate: f64,
118 pub learning_efficiency: f64,
120}
121
122#[derive(Debug, Clone, Serialize, Deserialize)]
124pub enum AdaptationUrgency {
125 Low,
126 Medium,
127 High,
128 Critical,
129}
130
131#[derive(Debug, Clone, Serialize, Deserialize)]
133pub struct GlobalLRInsights {
134 pub overall_efficiency: f64,
136 pub lr_distribution_health: f64,
138 pub gradient_flow_quality: f64,
140 pub training_stability: TrainingStability,
142 pub global_adjustments: Vec<GlobalLRAdjustment>,
144 pub critical_issues: Vec<String>,
146}
147
148#[derive(Debug, Clone, Serialize, Deserialize)]
150pub struct TrainingStability {
151 pub stability_score: f64,
153 pub instability_indicators: Vec<InstabilityIndicator>,
155 pub stability_trends: Vec<StabilityTrendPoint>,
157 pub predicted_stability: f64,
159}
160
161#[derive(Debug, Clone, Serialize, Deserialize)]
163pub struct InstabilityIndicator {
164 pub instability_type: InstabilityType,
166 pub severity: f64,
168 pub affected_layers: Vec<String>,
170 pub recommended_actions: Vec<String>,
172}
173
174#[derive(Debug, Clone, Serialize, Deserialize)]
176pub enum InstabilityType {
177 GradientExplosion,
178 GradientVanishing,
179 OscillatingLoss,
180 SlowConvergence,
181 WeightDivergence,
182 NumericalInstability,
183}
184
185#[derive(Debug, Clone, Serialize, Deserialize)]
187pub struct StabilityTrendPoint {
188 pub time_step: usize,
190 pub stability_score: f64,
192 pub contributing_factors: HashMap<String, f64>,
194}
195
196#[derive(Debug, Clone, Serialize, Deserialize)]
198pub struct GlobalLRAdjustment {
199 pub adjustment_type: GlobalAdjustmentType,
201 pub magnitude: f64,
203 pub expected_impact: f64,
205 pub priority: AdjustmentPriority,
207 pub instructions: String,
209}
210
211#[derive(Debug, Clone, Serialize, Deserialize)]
213pub enum GlobalAdjustmentType {
214 UniformScaling,
215 LayerTypeSpecific,
216 DepthDependent,
217 AdaptiveScheduling,
218 WarmupAdjustment,
219 DecayRateModification,
220}
221
222#[derive(Debug, Clone, Serialize, Deserialize)]
224pub enum AdjustmentPriority {
225 Low,
226 Medium,
227 High,
228 Immediate,
229}
230
231#[derive(Debug, Clone, Serialize, Deserialize)]
233pub struct LRAdaptationStrategy {
234 pub strategy_name: String,
236 pub description: String,
238 pub implementation_steps: Vec<ImplementationStep>,
240 pub expected_benefits: Vec<String>,
242 pub potential_risks: Vec<String>,
244 pub success_metrics: Vec<String>,
246 pub monitoring_requirements: Vec<String>,
248}
249
250#[derive(Debug, Clone, Serialize, Deserialize)]
252pub struct ImplementationStep {
253 pub step_number: usize,
255 pub description: String,
257 pub code_changes: Vec<String>,
259 pub timeline: String,
261 pub dependencies: Vec<String>,
263}
264
265#[derive(Debug, Clone, Serialize, Deserialize)]
267pub struct TrainingPhaseRecommendation {
268 pub phase_name: String,
270 pub duration_epochs: usize,
272 pub lr_schedule: LRSchedule,
274 pub objectives: Vec<String>,
276 pub success_criteria: Vec<String>,
278 pub transition_conditions: Vec<String>,
280}
281
282#[derive(Debug, Clone, Serialize, Deserialize)]
284pub struct LRSchedule {
285 pub schedule_type: LRScheduleType,
287 pub initial_lr: f64,
289 pub parameters: HashMap<String, f64>,
291 pub layer_multipliers: HashMap<String, f64>,
293}
294
295#[derive(Debug, Clone, Serialize, Deserialize)]
297pub enum LRScheduleType {
298 Constant,
299 LinearDecay,
300 ExponentialDecay,
301 CosineAnnealing,
302 StepDecay,
303 CyclicalLR,
304 OneCycleLR,
305 AdaptiveSchedule,
306}
307
308#[derive(Debug, Clone, Serialize, Deserialize)]
310pub struct LRSchedulePrediction {
311 pub schedule: LRSchedule,
313 pub predicted_accuracy: f64,
315 pub predicted_convergence_epochs: usize,
317 pub predicted_stability: f64,
319 pub prediction_confidence: f64,
321 pub risk_assessment: RiskAssessment,
323}
324
325#[derive(Debug, Clone, Serialize, Deserialize)]
327pub struct RiskAssessment {
328 pub overall_risk: RiskLevel,
330 pub specific_risks: Vec<SpecificRisk>,
332 pub mitigation_strategies: Vec<String>,
334}
335
336#[derive(Debug, Clone, Serialize, Deserialize)]
338pub enum RiskLevel {
339 VeryLow,
340 Low,
341 Medium,
342 High,
343 VeryHigh,
344}
345
346#[derive(Debug, Clone, Serialize, Deserialize)]
348pub struct SpecificRisk {
349 pub risk_type: String,
351 pub probability: f64,
353 pub impact: f64,
355 pub description: String,
357}
358
359#[derive(Debug, Clone, Serialize, Deserialize)]
361pub struct ModelSensitivityAnalysisResult {
362 pub timestamp: DateTime<Utc>,
364 pub hyperparameter_sensitivity: HyperparameterSensitivity,
366 pub architecture_sensitivity: ArchitectureSensitivity,
368 pub data_sensitivity: DataSensitivity,
370 pub training_sensitivity: TrainingSensitivity,
372 pub sensitivity_insights: SensitivityInsights,
374}
375
376#[derive(Debug, Clone, Serialize, Deserialize)]
378pub struct HyperparameterSensitivity {
379 pub learning_rate_sensitivity: ParameterSensitivity,
381 pub batch_size_sensitivity: ParameterSensitivity,
383 pub regularization_sensitivity: ParameterSensitivity,
385 pub architecture_param_sensitivity: HashMap<String, ParameterSensitivity>,
387 pub most_sensitive_params: Vec<String>,
389 pub least_sensitive_params: Vec<String>,
391 pub interaction_effects: Vec<ParameterInteraction>,
393}
394
395#[derive(Debug, Clone, Serialize, Deserialize)]
397pub struct ParameterSensitivity {
398 pub parameter_name: String,
400 pub current_value: f64,
402 pub sensitivity_score: f64,
404 pub optimal_range: (f64, f64),
406 pub impact_curve: Vec<(f64, f64)>,
408 pub stability_region: (f64, f64),
410 pub critical_thresholds: Vec<f64>,
412}
413
414#[derive(Debug, Clone, Serialize, Deserialize)]
416pub struct ParameterInteraction {
417 pub param1: String,
419 pub param2: String,
421 pub interaction_strength: f64,
423 pub interaction_type: InteractionType,
425 pub joint_optimal_region: HashMap<String, (f64, f64)>,
427}
428
429#[derive(Debug, Clone, Serialize, Deserialize)]
431pub enum InteractionType {
432 Synergistic,
433 Antagonistic,
434 Independent,
435 Conditional,
436}
437
438#[derive(Debug, Clone, Serialize, Deserialize)]
440pub struct ArchitectureSensitivity {
441 pub depth_sensitivity: ArchitecturalSensitivity,
443 pub width_sensitivity: ArchitecturalSensitivity,
445 pub attention_head_sensitivity: ArchitecturalSensitivity,
447 pub skip_connection_sensitivity: ArchitecturalSensitivity,
449 pub component_importance: HashMap<String, f64>,
451 pub bottlenecks: Vec<ArchitecturalBottleneck>,
453}
454
455#[derive(Debug, Clone, Serialize, Deserialize)]
457pub struct ArchitecturalSensitivity {
458 pub component_name: String,
460 pub change_sensitivity: f64,
462 pub degradation_curve: Vec<(f64, f64)>,
464 pub min_viable_config: f64,
466 pub optimal_config: f64,
468 pub diminishing_returns_threshold: f64,
470}
471
472#[derive(Debug, Clone, Serialize, Deserialize)]
474pub struct ArchitecturalBottleneck {
475 pub location: String,
477 pub bottleneck_type: BottleneckType,
479 pub severity: f64,
481 pub performance_impact: f64,
483 pub resolution_recommendations: Vec<String>,
485}
486
487#[derive(Debug, Clone, Serialize, Deserialize)]
489pub enum BottleneckType {
490 ComputationalBottleneck,
491 MemoryBottleneck,
492 InformationBottleneck,
493 CapacityBottleneck,
494 CommunicationBottleneck,
495}
496
497#[derive(Debug, Clone, Serialize, Deserialize)]
499pub struct DataSensitivity {
500 pub data_size_sensitivity: DataSizeSensitivity,
502 pub data_quality_sensitivity: DataQualitySensitivity,
504 pub distribution_sensitivity: DistributionSensitivity,
506 pub feature_sensitivity: FeatureSensitivityAnalysis,
508}
509
510#[derive(Debug, Clone, Serialize, Deserialize)]
512pub struct DataSizeSensitivity {
513 pub current_size: usize,
515 pub minimum_effective_size: usize,
517 pub performance_curve: Vec<(usize, f64)>,
519 pub data_efficiency: f64,
521 pub diminishing_returns_point: usize,
523}
524
525#[derive(Debug, Clone, Serialize, Deserialize)]
527pub struct DataQualitySensitivity {
528 pub noise_tolerance: f64,
530 pub label_quality_importance: f64,
532 pub feature_quality_importance: f64,
534 pub quality_impact_curve: Vec<(f64, f64)>,
536}
537
538#[derive(Debug, Clone, Serialize, Deserialize)]
540pub struct DistributionSensitivity {
541 pub shift_sensitivity: f64,
543 pub imbalance_sensitivity: f64,
545 pub domain_adaptation_requirements: Vec<String>,
547 pub distribution_robustness: f64,
549}
550
551#[derive(Debug, Clone, Serialize, Deserialize)]
553pub struct FeatureSensitivityAnalysis {
554 pub most_important_features: Vec<String>,
556 pub least_important_features: Vec<String>,
558 pub feature_interactions: HashMap<(String, String), f64>,
560 pub feature_stability: HashMap<String, f64>,
562}
563
564#[derive(Debug, Clone, Serialize, Deserialize)]
566pub struct TrainingSensitivity {
567 pub initialization_sensitivity: InitializationSensitivity,
569 pub optimization_sensitivity: OptimizationSensitivity,
571 pub schedule_sensitivity: ScheduleSensitivity,
573 pub regularization_sensitivity: RegularizationSensitivity,
575}
576
577#[derive(Debug, Clone, Serialize, Deserialize)]
579pub struct InitializationSensitivity {
580 pub weight_init_sensitivity: f64,
582 pub bias_init_sensitivity: f64,
584 pub seed_sensitivity: f64,
586 pub scheme_importance: HashMap<String, f64>,
588}
589
590#[derive(Debug, Clone, Serialize, Deserialize)]
592pub struct OptimizationSensitivity {
593 pub optimizer_sensitivity: f64,
595 pub momentum_sensitivity: f64,
597 pub second_moment_sensitivity: f64,
599 pub optimizer_comparison: HashMap<String, f64>,
601}
602
603#[derive(Debug, Clone, Serialize, Deserialize)]
605pub struct ScheduleSensitivity {
606 pub lr_schedule_sensitivity: f64,
608 pub duration_sensitivity: f64,
610 pub warmup_sensitivity: f64,
612 pub schedule_param_importance: HashMap<String, f64>,
614}
615
616#[derive(Debug, Clone, Serialize, Deserialize)]
618pub struct RegularizationSensitivity {
619 pub dropout_sensitivity: f64,
621 pub weight_decay_sensitivity: f64,
623 pub batch_norm_sensitivity: f64,
625 pub method_comparison: HashMap<String, f64>,
627}
628
629#[derive(Debug, Clone, Serialize, Deserialize)]
631pub struct SensitivityInsights {
632 pub most_critical_factors: Vec<String>,
634 pub least_critical_factors: Vec<String>,
636 pub surprising_findings: Vec<String>,
638 pub robustness_assessment: RobustnessAssessment,
640 pub optimization_recommendations: Vec<String>,
642}
643
644#[derive(Debug, Clone, Serialize, Deserialize)]
646pub struct RobustnessAssessment {
647 pub overall_robustness: f64,
649 pub category_robustness: HashMap<String, f64>,
651 pub vulnerabilities: Vec<Vulnerability>,
653 pub strengths: Vec<String>,
655}
656
657#[derive(Debug, Clone, Serialize, Deserialize)]
659pub struct Vulnerability {
660 pub vulnerability_type: String,
662 pub severity: f64,
664 pub impact: String,
666 pub mitigation_strategies: Vec<String>,
668}
669
670#[derive(Debug)]
672pub struct AdvancedMLDebugger {
673 config: AdvancedMLDebuggingConfig,
674 lr_analysis_results: Vec<LayerWiseLRAnalysisResult>,
675 sensitivity_analysis_results: Vec<ModelSensitivityAnalysisResult>,
676}
677
678impl AdvancedMLDebugger {
679 pub fn new(config: AdvancedMLDebuggingConfig) -> Self {
681 Self {
682 config,
683 lr_analysis_results: Vec::new(),
684 sensitivity_analysis_results: Vec::new(),
685 }
686 }
687
688 pub async fn analyze_layer_wise_learning_rates(
690 &mut self,
691 layer_gradients: &HashMap<String, ArrayD<f32>>,
692 layer_weights: &HashMap<String, ArrayD<f32>>,
693 current_lr: f64,
694 loss_history: &[f64],
695 ) -> Result<LayerWiseLRAnalysisResult> {
696 if !self.config.enable_layer_wise_lr_analysis {
697 return Err(anyhow::anyhow!(
698 "Layer-wise learning rate analysis is disabled"
699 ));
700 }
701
702 let mut layer_lr_recommendations = HashMap::new();
703
704 for (layer_id, gradients) in layer_gradients {
706 if let Some(weights) = layer_weights.get(layer_id) {
707 let recommendation = self.analyze_single_layer_lr(
708 layer_id,
709 gradients,
710 weights,
711 current_lr,
712 loss_history,
713 );
714 layer_lr_recommendations.insert(layer_id.clone(), recommendation);
715 }
716 }
717
718 let global_lr_insights =
720 self.generate_global_lr_insights(&layer_lr_recommendations, loss_history);
721
722 let adaptation_strategy =
724 self.create_lr_adaptation_strategy(&layer_lr_recommendations, &global_lr_insights);
725
726 let training_phase_recommendations =
728 self.generate_training_phase_recommendations(&adaptation_strategy);
729
730 let lr_schedule_predictions =
732 self.predict_lr_schedule_performance(&layer_lr_recommendations);
733
734 let result = LayerWiseLRAnalysisResult {
735 timestamp: Utc::now(),
736 layer_lr_recommendations,
737 global_lr_insights,
738 adaptation_strategy,
739 training_phase_recommendations,
740 lr_schedule_predictions,
741 };
742
743 self.lr_analysis_results.push(result.clone());
744 Ok(result)
745 }
746
747 pub async fn analyze_model_sensitivity(
768 &mut self,
769 model_params: &HashMap<String, f64>,
770 performance_metrics: &[f64],
771 architecture_config: &HashMap<String, f64>,
772 ) -> Result<ModelSensitivityAnalysisResult> {
773 if !self.config.enable_model_sensitivity_analysis {
774 return Err(anyhow::anyhow!("Model sensitivity analysis is disabled"));
775 }
776 Err(anyhow::anyhow!(
777 "model sensitivity analysis is not implemented: it needs the model re-evaluated at \
778 perturbed settings, but this call receives only {} current parameter values, {} \
779 architecture values and {} already-observed metric samples -- no perturbation \
780 response is derivable from a single operating point",
781 model_params.len(),
782 architecture_config.len(),
783 performance_metrics.len(),
784 ))
785 }
786
787 pub async fn generate_report(&self) -> Result<AdvancedMLDebuggingReport> {
789 Ok(AdvancedMLDebuggingReport {
790 timestamp: Utc::now(),
791 config: self.config.clone(),
792 lr_analysis_count: self.lr_analysis_results.len(),
793 sensitivity_analysis_count: self.sensitivity_analysis_results.len(),
794 recent_lr_analyses: self.lr_analysis_results.iter().rev().take(3).cloned().collect(),
795 recent_sensitivity_analyses: self
796 .sensitivity_analysis_results
797 .iter()
798 .rev()
799 .take(3)
800 .cloned()
801 .collect(),
802 advanced_insights: self.generate_advanced_insights(),
803 })
804 }
805
806 fn analyze_single_layer_lr(
809 &self,
810 layer_id: &str,
811 gradients: &ArrayD<f32>,
812 weights: &ArrayD<f32>,
813 current_lr: f64,
814 loss_history: &[f64],
815 ) -> LayerLRRecommendation {
816 let gradient_magnitude =
818 gradients.iter().map(|&x| x.abs() as f64).sum::<f64>() / gradients.len() as f64;
819 let weight_magnitude =
820 weights.iter().map(|&x| x.abs() as f64).sum::<f64>() / weights.len() as f64;
821
822 let gradient_variance =
824 gradients.iter().map(|&x| (x as f64 - gradient_magnitude).powi(2)).sum::<f64>()
825 / gradients.len() as f64;
826
827 let gradient_norm = gradient_magnitude;
828 let recommended_lr = if gradient_norm > 0.0 {
829 let base_lr = 0.001;
831 let adaptation_factor = (1.0 / (1.0 + gradient_variance)).sqrt();
832 let magnitude_factor = (1.0 / (1.0 + gradient_norm)).sqrt();
833 base_lr * adaptation_factor * magnitude_factor * 10.0
834 } else {
835 current_lr
836 };
837
838 let layer_metrics = LayerLRMetrics {
840 gradient_magnitude,
841 weight_update_magnitude: gradient_magnitude * current_lr,
842 parameter_norm: weight_magnitude,
843 loss_contribution: self.estimate_layer_loss_contribution(loss_history),
844 stability_score: self.calculate_layer_stability(gradients),
845 convergence_rate: self.estimate_convergence_rate(loss_history),
846 learning_efficiency: gradient_magnitude / (weight_magnitude + 1e-8),
847 };
848
849 let lr_ratio = recommended_lr / current_lr;
851 let urgency = if !(0.1..=10.0).contains(&lr_ratio) {
852 AdaptationUrgency::Critical
853 } else if !(0.33..=3.0).contains(&lr_ratio) {
854 AdaptationUrgency::High
855 } else if !(0.67..=1.5).contains(&lr_ratio) {
856 AdaptationUrgency::Medium
857 } else {
858 AdaptationUrgency::Low
859 };
860
861 let reasoning = if recommended_lr > current_lr * 1.2 {
863 "Layer shows slow learning with small gradients, increase learning rate".to_string()
864 } else if recommended_lr < current_lr * 0.8 {
865 "Layer shows instability or large gradients, decrease learning rate".to_string()
866 } else {
867 "Current learning rate appears appropriate for this layer".to_string()
868 };
869
870 LayerLRRecommendation {
871 layer_id: layer_id.to_string(),
872 layer_type: self.infer_layer_type(layer_id),
873 current_lr,
874 recommended_lr,
875 confidence: None,
880 reasoning,
881 layer_metrics,
882 lr_sensitivity: lr_ratio.abs(),
883 urgency,
884 }
885 }
886
887 fn generate_global_lr_insights(
888 &self,
889 layer_recommendations: &HashMap<String, LayerLRRecommendation>,
890 loss_history: &[f64],
891 ) -> GlobalLRInsights {
892 let overall_efficiency = layer_recommendations
893 .values()
894 .map(|rec| rec.layer_metrics.learning_efficiency)
895 .sum::<f64>()
896 / layer_recommendations.len() as f64;
897
898 let lr_distribution_health = self.calculate_lr_distribution_health(layer_recommendations);
899 let gradient_flow_quality = self.calculate_gradient_flow_quality(layer_recommendations);
900 let training_stability =
901 self.assess_training_stability(layer_recommendations, loss_history);
902 let global_adjustments = self.generate_global_adjustments(layer_recommendations);
903 let critical_issues = self.identify_critical_issues(layer_recommendations);
904
905 GlobalLRInsights {
906 overall_efficiency,
907 lr_distribution_health,
908 gradient_flow_quality,
909 training_stability,
910 global_adjustments,
911 critical_issues,
912 }
913 }
914
915 fn create_lr_adaptation_strategy(
916 &self,
917 _layer_recommendations: &HashMap<String, LayerLRRecommendation>,
918 global_insights: &GlobalLRInsights,
919 ) -> LRAdaptationStrategy {
920 let strategy_name = if global_insights.overall_efficiency < 0.5 {
922 "Aggressive Learning Rate Adaptation".to_string()
923 } else {
924 "Conservative Learning Rate Tuning".to_string()
925 };
926
927 LRAdaptationStrategy {
928 strategy_name: strategy_name.clone(),
929 description: "Strategy to optimize learning rates based on current model state"
930 .to_string(),
931 implementation_steps: vec![ImplementationStep {
932 step_number: 1,
933 description: "Implement layer-wise learning rate multipliers".to_string(),
934 code_changes: vec!["Add lr_multipliers to optimizer config".to_string()],
935 timeline: "1-2 days".to_string(),
936 dependencies: vec!["Optimizer modification".to_string()],
937 }],
938 expected_benefits: vec![
939 "Improved convergence speed".to_string(),
940 "Better training stability".to_string(),
941 "Reduced overfitting risk".to_string(),
942 ],
943 potential_risks: vec!["Initial instability during adaptation".to_string()],
944 success_metrics: vec![
945 "Faster loss reduction".to_string(),
946 "Improved validation accuracy".to_string(),
947 ],
948 monitoring_requirements: vec!["Track per-layer gradient norms".to_string()],
949 }
950 }
951
952 fn generate_training_phase_recommendations(
953 &self,
954 _strategy: &LRAdaptationStrategy,
955 ) -> Vec<TrainingPhaseRecommendation> {
956 vec![TrainingPhaseRecommendation {
957 phase_name: "Warmup Phase".to_string(),
958 duration_epochs: 5,
959 lr_schedule: LRSchedule {
960 schedule_type: LRScheduleType::LinearDecay,
961 initial_lr: 0.0001,
962 parameters: HashMap::new(),
963 layer_multipliers: HashMap::new(),
964 },
965 objectives: vec!["Stabilize training".to_string()],
966 success_criteria: vec!["Decreasing loss".to_string()],
967 transition_conditions: vec!["Stable gradient norms".to_string()],
968 }]
969 }
970
971 fn predict_lr_schedule_performance(
972 &self,
973 _layer_recommendations: &HashMap<String, LayerLRRecommendation>,
974 ) -> Vec<LRSchedulePrediction> {
975 vec![LRSchedulePrediction {
976 schedule: LRSchedule {
977 schedule_type: LRScheduleType::ExponentialDecay,
978 initial_lr: 0.001,
979 parameters: HashMap::new(),
980 layer_multipliers: HashMap::new(),
981 },
982 predicted_accuracy: 0.92,
983 predicted_convergence_epochs: 50,
984 predicted_stability: 0.8,
985 prediction_confidence: 0.7,
986 risk_assessment: RiskAssessment {
987 overall_risk: RiskLevel::Medium,
988 specific_risks: vec![],
989 mitigation_strategies: vec![],
990 },
991 }]
992 }
993
994 fn estimate_layer_loss_contribution(&self, loss_history: &[f64]) -> f64 {
999 if loss_history.len() >= 2 {
1001 (loss_history[loss_history.len() - 2] - loss_history[loss_history.len() - 1]).abs()
1002 } else {
1003 0.1
1004 }
1005 }
1006
1007 fn calculate_layer_stability(&self, gradients: &ArrayD<f32>) -> f64 {
1008 let gradient_variance = gradients.iter().map(|&x| x as f64).collect::<Vec<_>>();
1009
1010 if gradient_variance.is_empty() {
1011 return 0.5;
1012 }
1013
1014 let mean = gradient_variance.iter().sum::<f64>() / gradient_variance.len() as f64;
1015 let variance = gradient_variance.iter().map(|&x| (x - mean).powi(2)).sum::<f64>()
1016 / gradient_variance.len() as f64;
1017
1018 1.0 / (1.0 + variance) }
1020
1021 fn estimate_convergence_rate(&self, loss_history: &[f64]) -> f64 {
1022 if loss_history.len() < 3 {
1023 return 0.5;
1024 }
1025
1026 let recent_improvement =
1027 loss_history[loss_history.len() - 3] - loss_history[loss_history.len() - 1];
1028 recent_improvement.abs()
1029 }
1030
1031 fn infer_layer_type(&self, layer_id: &str) -> String {
1032 if layer_id.contains("attention") {
1033 "attention".to_string()
1034 } else if layer_id.contains("feedforward") || layer_id.contains("mlp") {
1035 "feedforward".to_string()
1036 } else if layer_id.contains("embedding") {
1037 "embedding".to_string()
1038 } else {
1039 "unknown".to_string()
1040 }
1041 }
1042
1043 fn calculate_lr_distribution_health(
1044 &self,
1045 recommendations: &HashMap<String, LayerLRRecommendation>,
1046 ) -> f64 {
1047 let lr_ratios: Vec<f64> = recommendations
1048 .values()
1049 .map(|rec| rec.recommended_lr / rec.current_lr)
1050 .collect();
1051
1052 if lr_ratios.is_empty() {
1053 return 0.5;
1054 }
1055
1056 let mean_ratio = lr_ratios.iter().sum::<f64>() / lr_ratios.len() as f64;
1057 let variance = lr_ratios.iter().map(|&x| (x - mean_ratio).powi(2)).sum::<f64>()
1058 / lr_ratios.len() as f64;
1059
1060 1.0 / (1.0 + variance) }
1062
1063 fn calculate_gradient_flow_quality(
1064 &self,
1065 recommendations: &HashMap<String, LayerLRRecommendation>,
1066 ) -> f64 {
1067 recommendations
1068 .values()
1069 .map(|rec| rec.layer_metrics.stability_score)
1070 .sum::<f64>()
1071 / recommendations.len() as f64
1072 }
1073
1074 fn assess_training_stability(
1075 &self,
1076 recommendations: &HashMap<String, LayerLRRecommendation>,
1077 _loss_history: &[f64],
1078 ) -> TrainingStability {
1079 let stability_score = recommendations
1080 .values()
1081 .map(|rec| rec.layer_metrics.stability_score)
1082 .sum::<f64>()
1083 / recommendations.len() as f64;
1084
1085 TrainingStability {
1086 stability_score,
1087 instability_indicators: vec![],
1088 stability_trends: vec![],
1089 predicted_stability: stability_score * 0.9, }
1091 }
1092
1093 fn generate_global_adjustments(
1094 &self,
1095 _recommendations: &HashMap<String, LayerLRRecommendation>,
1096 ) -> Vec<GlobalLRAdjustment> {
1097 vec![GlobalLRAdjustment {
1098 adjustment_type: GlobalAdjustmentType::LayerTypeSpecific,
1099 magnitude: 1.5,
1100 expected_impact: 0.1,
1101 priority: AdjustmentPriority::Medium,
1102 instructions: "Apply different learning rates to attention vs feedforward layers"
1103 .to_string(),
1104 }]
1105 }
1106
1107 fn identify_critical_issues(
1108 &self,
1109 recommendations: &HashMap<String, LayerLRRecommendation>,
1110 ) -> Vec<String> {
1111 let mut issues = Vec::new();
1112
1113 for recommendation in recommendations.values() {
1114 if matches!(recommendation.urgency, AdaptationUrgency::Critical) {
1115 issues.push(format!(
1116 "Critical learning rate issue in layer {}",
1117 recommendation.layer_id
1118 ));
1119 }
1120 }
1121
1122 issues
1123 }
1124
1125 fn generate_advanced_insights(&self) -> HashMap<String, String> {
1126 let mut insights = HashMap::new();
1127
1128 insights.insert(
1129 "total_lr_analyses".to_string(),
1130 self.lr_analysis_results.len().to_string(),
1131 );
1132 insights.insert(
1133 "total_sensitivity_analyses".to_string(),
1134 self.sensitivity_analysis_results.len().to_string(),
1135 );
1136
1137 if let Some(latest_lr) = self.lr_analysis_results.last() {
1138 insights.insert(
1139 "latest_lr_efficiency".to_string(),
1140 format!("{:.2}", latest_lr.global_lr_insights.overall_efficiency),
1141 );
1142 }
1143
1144 insights
1145 }
1146}
1147
1148#[derive(Debug, Clone, Serialize, Deserialize)]
1150pub struct AdvancedMLDebuggingReport {
1151 pub timestamp: DateTime<Utc>,
1152 pub config: AdvancedMLDebuggingConfig,
1153 pub lr_analysis_count: usize,
1154 pub sensitivity_analysis_count: usize,
1155 pub recent_lr_analyses: Vec<LayerWiseLRAnalysisResult>,
1156 pub recent_sensitivity_analyses: Vec<ModelSensitivityAnalysisResult>,
1157 pub advanced_insights: HashMap<String, String>,
1158}
1159
1160#[cfg(test)]
1161mod tests {
1162 use super::*;
1163
1164 #[tokio::test]
1167 async fn model_sensitivity_analysis_refuses_instead_of_inventing_curves() {
1168 let config = AdvancedMLDebuggingConfig::default();
1169 let mut debugger = AdvancedMLDebugger::new(config);
1170 let mut params = HashMap::new();
1171 params.insert("learning_rate".to_string(), 0.001);
1172 let err = debugger
1173 .analyze_model_sensitivity(¶ms, &[0.9, 0.91, 0.92], &HashMap::new())
1174 .await
1175 .expect_err("a single operating point cannot yield a sensitivity analysis");
1176 let msg = err.to_string();
1177 assert!(
1178 msg.contains("not implemented"),
1179 "must name what is missing: {msg}"
1180 );
1181 assert!(
1182 msg.contains("perturbation response"),
1183 "must explain why it is unmeasurable: {msg}"
1184 );
1185 }
1186
1187 #[tokio::test]
1188 async fn layer_lr_recommendation_confidence_is_absent_not_a_constant() {
1189 let config = AdvancedMLDebuggingConfig::default();
1190 let mut debugger = AdvancedMLDebugger::new(config);
1191 let mut layer_gradients = HashMap::new();
1192 let mut layer_weights = HashMap::new();
1193 let gradients = ArrayD::from_shape_vec(
1194 scirs2_core::ndarray::IxDyn(&[4]),
1195 vec![0.1f32, -0.2, 0.3, -0.4],
1196 )
1197 .expect("gradient tensor");
1198 let weights = ArrayD::from_shape_vec(
1199 scirs2_core::ndarray::IxDyn(&[4]),
1200 vec![1.0f32, 2.0, 3.0, 4.0],
1201 )
1202 .expect("weight tensor");
1203 layer_gradients.insert("layer0".to_string(), gradients);
1204 layer_weights.insert("layer0".to_string(), weights);
1205
1206 let result = debugger
1207 .analyze_layer_wise_learning_rates(
1208 &layer_gradients,
1209 &layer_weights,
1210 0.001,
1211 &[1.0, 0.9, 0.8],
1212 )
1213 .await
1214 .expect("analysis should run");
1215 for rec in result.layer_lr_recommendations.values() {
1216 assert_eq!(
1217 rec.confidence, None,
1218 "a one-snapshot recommendation has no statistical confidence (was 0.8)"
1219 );
1220 }
1221 }
1222
1223 #[tokio::test]
1224 async fn test_advanced_ml_debugger_creation() {
1225 let config = AdvancedMLDebuggingConfig::default();
1226 let debugger = AdvancedMLDebugger::new(config);
1227 assert_eq!(debugger.lr_analysis_results.len(), 0);
1228 }
1229
1230 #[tokio::test]
1231 async fn test_layer_wise_lr_analysis() {
1232 let config = AdvancedMLDebuggingConfig::default();
1233 let mut debugger = AdvancedMLDebugger::new(config);
1234
1235 let mut layer_gradients = HashMap::new();
1236 let mut layer_weights = HashMap::new();
1237
1238 let gradients =
1240 ArrayD::from_shape_vec(vec![10, 10], (0..100).map(|x| x as f32 * 0.01).collect())
1241 .expect("operation failed in test");
1242 let weights =
1243 ArrayD::from_shape_vec(vec![10, 10], (0..100).map(|x| x as f32 * 0.1).collect())
1244 .expect("operation failed in test");
1245
1246 layer_gradients.insert("layer_0".to_string(), gradients);
1247 layer_weights.insert("layer_0".to_string(), weights);
1248
1249 let loss_history = vec![1.0, 0.8, 0.6, 0.5];
1250
1251 let result = debugger
1252 .analyze_layer_wise_learning_rates(
1253 &layer_gradients,
1254 &layer_weights,
1255 0.001,
1256 &loss_history,
1257 )
1258 .await;
1259 assert!(result.is_ok());
1260
1261 let analysis = result.expect("operation failed in test");
1262 assert_eq!(analysis.layer_lr_recommendations.len(), 1);
1263 assert!(analysis.layer_lr_recommendations.contains_key("layer_0"));
1264 }
1265}