Skip to main content

trustformers_debug/
advanced_ml_debugging.rs

1//! # Advanced ML Debugging Tools
2//!
3//! Advanced machine learning specific debugging techniques including layer-wise learning rate adaptation,
4//! model sensitivity analysis, gradient flow optimization, and neural architecture debugging.
5
6use anyhow::Result;
7use chrono::{DateTime, Utc};
8use scirs2_core::ndarray::*; // SciRS2 Integration Policy - was: use ndarray::{Array1, Array2, Array3, ArrayD};
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11
12/// Configuration for advanced ML debugging
13#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct AdvancedMLDebuggingConfig {
15    /// Enable layer-wise learning rate analysis
16    pub enable_layer_wise_lr_analysis: bool,
17    /// Enable model sensitivity analysis
18    pub enable_model_sensitivity_analysis: bool,
19    /// Enable gradient flow optimization analysis
20    pub enable_gradient_flow_optimization: bool,
21    /// Enable neural architecture debugging
22    pub enable_neural_architecture_debugging: bool,
23    /// Enable activation pattern analysis
24    pub enable_activation_pattern_analysis: bool,
25    /// Enable weight distribution analysis
26    pub enable_weight_distribution_analysis: bool,
27    /// Enable training dynamics analysis
28    pub enable_training_dynamics_analysis: bool,
29    /// Enable optimization landscape analysis
30    pub enable_optimization_landscape_analysis: bool,
31    /// Number of samples for sensitivity analysis
32    pub sensitivity_samples: usize,
33    /// Learning rate adaptation threshold
34    pub lr_adaptation_threshold: f64,
35    /// Maximum number of layers to analyze
36    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/// Layer-wise learning rate adaptation analysis result
58#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct LayerWiseLRAnalysisResult {
60    /// Analysis timestamp
61    pub timestamp: DateTime<Utc>,
62    /// Learning rate recommendations per layer
63    pub layer_lr_recommendations: HashMap<String, LayerLRRecommendation>,
64    /// Global learning rate insights
65    pub global_lr_insights: GlobalLRInsights,
66    /// Learning rate adaptation strategy
67    pub adaptation_strategy: LRAdaptationStrategy,
68    /// Training phase recommendations
69    pub training_phase_recommendations: Vec<TrainingPhaseRecommendation>,
70    /// Performance predictions with different LR schedules
71    pub lr_schedule_predictions: Vec<LRSchedulePrediction>,
72}
73
74/// Learning rate recommendation for a specific layer
75#[derive(Debug, Clone, Serialize, Deserialize)]
76pub struct LayerLRRecommendation {
77    /// Layer identifier
78    pub layer_id: String,
79    /// Layer type (e.g., "attention", "feedforward", "embedding")
80    pub layer_type: String,
81    /// Current learning rate
82    pub current_lr: f64,
83    /// Recommended learning rate
84    pub recommended_lr: f64,
85    /// Statistical confidence in the recommendation.
86    ///
87    /// `None` from [`AdvancedMLDebugger::analyze_layer_wise_learning_rates`]:
88    /// that path sees a single gradient/weight snapshot per layer, with no
89    /// repeated trials or held-out evaluation, so there is no sampling
90    /// distribution to derive a confidence from. It used to be a hardcoded
91    /// `0.8`.
92    pub confidence: Option<f64>,
93    /// Reasoning for recommendation
94    pub reasoning: String,
95    /// Layer-specific metrics
96    pub layer_metrics: LayerLRMetrics,
97    /// Sensitivity to learning rate changes
98    pub lr_sensitivity: f64,
99    /// Adaptation urgency level
100    pub urgency: AdaptationUrgency,
101}
102
103/// Layer-specific learning rate metrics
104#[derive(Debug, Clone, Serialize, Deserialize)]
105pub struct LayerLRMetrics {
106    /// Gradient magnitude
107    pub gradient_magnitude: f64,
108    /// Weight update magnitude
109    pub weight_update_magnitude: f64,
110    /// Parameter norm
111    pub parameter_norm: f64,
112    /// Loss contribution
113    pub loss_contribution: f64,
114    /// Training stability score
115    pub stability_score: f64,
116    /// Convergence rate
117    pub convergence_rate: f64,
118    /// Learning efficiency
119    pub learning_efficiency: f64,
120}
121
122/// Urgency level for learning rate adaptation
123#[derive(Debug, Clone, Serialize, Deserialize)]
124pub enum AdaptationUrgency {
125    Low,
126    Medium,
127    High,
128    Critical,
129}
130
131/// Global learning rate insights
132#[derive(Debug, Clone, Serialize, Deserialize)]
133pub struct GlobalLRInsights {
134    /// Overall model learning efficiency
135    pub overall_efficiency: f64,
136    /// Learning rate distribution health
137    pub lr_distribution_health: f64,
138    /// Gradient flow quality
139    pub gradient_flow_quality: f64,
140    /// Training stability assessment
141    pub training_stability: TrainingStability,
142    /// Recommended global adjustments
143    pub global_adjustments: Vec<GlobalLRAdjustment>,
144    /// Critical issues requiring immediate attention
145    pub critical_issues: Vec<String>,
146}
147
148/// Training stability assessment
149#[derive(Debug, Clone, Serialize, Deserialize)]
150pub struct TrainingStability {
151    /// Stability score (0-1)
152    pub stability_score: f64,
153    /// Instability indicators
154    pub instability_indicators: Vec<InstabilityIndicator>,
155    /// Stability trends over time
156    pub stability_trends: Vec<StabilityTrendPoint>,
157    /// Predicted stability with current settings
158    pub predicted_stability: f64,
159}
160
161/// Indicator of training instability
162#[derive(Debug, Clone, Serialize, Deserialize)]
163pub struct InstabilityIndicator {
164    /// Type of instability
165    pub instability_type: InstabilityType,
166    /// Severity level
167    pub severity: f64,
168    /// Affected layers
169    pub affected_layers: Vec<String>,
170    /// Recommended actions
171    pub recommended_actions: Vec<String>,
172}
173
174/// Type of training instability
175#[derive(Debug, Clone, Serialize, Deserialize)]
176pub enum InstabilityType {
177    GradientExplosion,
178    GradientVanishing,
179    OscillatingLoss,
180    SlowConvergence,
181    WeightDivergence,
182    NumericalInstability,
183}
184
185/// Point in stability trend analysis
186#[derive(Debug, Clone, Serialize, Deserialize)]
187pub struct StabilityTrendPoint {
188    /// Time step or epoch
189    pub time_step: usize,
190    /// Stability score at this point
191    pub stability_score: f64,
192    /// Contributing factors
193    pub contributing_factors: HashMap<String, f64>,
194}
195
196/// Global learning rate adjustment recommendation
197#[derive(Debug, Clone, Serialize, Deserialize)]
198pub struct GlobalLRAdjustment {
199    /// Adjustment type
200    pub adjustment_type: GlobalAdjustmentType,
201    /// Adjustment magnitude
202    pub magnitude: f64,
203    /// Expected impact
204    pub expected_impact: f64,
205    /// Implementation priority
206    pub priority: AdjustmentPriority,
207    /// Implementation instructions
208    pub instructions: String,
209}
210
211/// Type of global learning rate adjustment
212#[derive(Debug, Clone, Serialize, Deserialize)]
213pub enum GlobalAdjustmentType {
214    UniformScaling,
215    LayerTypeSpecific,
216    DepthDependent,
217    AdaptiveScheduling,
218    WarmupAdjustment,
219    DecayRateModification,
220}
221
222/// Priority level for adjustments
223#[derive(Debug, Clone, Serialize, Deserialize)]
224pub enum AdjustmentPriority {
225    Low,
226    Medium,
227    High,
228    Immediate,
229}
230
231/// Learning rate adaptation strategy
232#[derive(Debug, Clone, Serialize, Deserialize)]
233pub struct LRAdaptationStrategy {
234    /// Strategy name
235    pub strategy_name: String,
236    /// Strategy description
237    pub description: String,
238    /// Implementation steps
239    pub implementation_steps: Vec<ImplementationStep>,
240    /// Expected benefits
241    pub expected_benefits: Vec<String>,
242    /// Potential risks
243    pub potential_risks: Vec<String>,
244    /// Success metrics
245    pub success_metrics: Vec<String>,
246    /// Monitoring requirements
247    pub monitoring_requirements: Vec<String>,
248}
249
250/// Step in implementing an adaptation strategy
251#[derive(Debug, Clone, Serialize, Deserialize)]
252pub struct ImplementationStep {
253    /// Step number
254    pub step_number: usize,
255    /// Step description
256    pub description: String,
257    /// Code changes required
258    pub code_changes: Vec<String>,
259    /// Expected timeline
260    pub timeline: String,
261    /// Dependencies
262    pub dependencies: Vec<String>,
263}
264
265/// Training phase recommendation
266#[derive(Debug, Clone, Serialize, Deserialize)]
267pub struct TrainingPhaseRecommendation {
268    /// Phase name
269    pub phase_name: String,
270    /// Phase duration (epochs)
271    pub duration_epochs: usize,
272    /// Learning rate schedule for this phase
273    pub lr_schedule: LRSchedule,
274    /// Phase objectives
275    pub objectives: Vec<String>,
276    /// Success criteria
277    pub success_criteria: Vec<String>,
278    /// Transition conditions
279    pub transition_conditions: Vec<String>,
280}
281
282/// Learning rate schedule definition
283#[derive(Debug, Clone, Serialize, Deserialize)]
284pub struct LRSchedule {
285    /// Schedule type
286    pub schedule_type: LRScheduleType,
287    /// Initial learning rate
288    pub initial_lr: f64,
289    /// Schedule parameters
290    pub parameters: HashMap<String, f64>,
291    /// Layer-specific multipliers
292    pub layer_multipliers: HashMap<String, f64>,
293}
294
295/// Type of learning rate schedule
296#[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/// Prediction of performance with different LR schedules
309#[derive(Debug, Clone, Serialize, Deserialize)]
310pub struct LRSchedulePrediction {
311    /// Schedule being evaluated
312    pub schedule: LRSchedule,
313    /// Predicted final accuracy
314    pub predicted_accuracy: f64,
315    /// Predicted convergence time
316    pub predicted_convergence_epochs: usize,
317    /// Predicted training stability
318    pub predicted_stability: f64,
319    /// Confidence in prediction
320    pub prediction_confidence: f64,
321    /// Risk assessment
322    pub risk_assessment: RiskAssessment,
323}
324
325/// Risk assessment for a learning rate schedule
326#[derive(Debug, Clone, Serialize, Deserialize)]
327pub struct RiskAssessment {
328    /// Overall risk level
329    pub overall_risk: RiskLevel,
330    /// Specific risks
331    pub specific_risks: Vec<SpecificRisk>,
332    /// Mitigation strategies
333    pub mitigation_strategies: Vec<String>,
334}
335
336/// Risk level assessment
337#[derive(Debug, Clone, Serialize, Deserialize)]
338pub enum RiskLevel {
339    VeryLow,
340    Low,
341    Medium,
342    High,
343    VeryHigh,
344}
345
346/// Specific risk in training
347#[derive(Debug, Clone, Serialize, Deserialize)]
348pub struct SpecificRisk {
349    /// Risk type
350    pub risk_type: String,
351    /// Probability of occurrence
352    pub probability: f64,
353    /// Impact severity
354    pub impact: f64,
355    /// Description
356    pub description: String,
357}
358
359/// Model sensitivity analysis result
360#[derive(Debug, Clone, Serialize, Deserialize)]
361pub struct ModelSensitivityAnalysisResult {
362    /// Analysis timestamp
363    pub timestamp: DateTime<Utc>,
364    /// Hyperparameter sensitivity analysis
365    pub hyperparameter_sensitivity: HyperparameterSensitivity,
366    /// Architecture sensitivity analysis
367    pub architecture_sensitivity: ArchitectureSensitivity,
368    /// Data sensitivity analysis
369    pub data_sensitivity: DataSensitivity,
370    /// Training procedure sensitivity
371    pub training_sensitivity: TrainingSensitivity,
372    /// Overall sensitivity insights
373    pub sensitivity_insights: SensitivityInsights,
374}
375
376/// Hyperparameter sensitivity analysis
377#[derive(Debug, Clone, Serialize, Deserialize)]
378pub struct HyperparameterSensitivity {
379    /// Learning rate sensitivity
380    pub learning_rate_sensitivity: ParameterSensitivity,
381    /// Batch size sensitivity
382    pub batch_size_sensitivity: ParameterSensitivity,
383    /// Regularization sensitivity
384    pub regularization_sensitivity: ParameterSensitivity,
385    /// Architecture parameter sensitivity
386    pub architecture_param_sensitivity: HashMap<String, ParameterSensitivity>,
387    /// Most sensitive parameters
388    pub most_sensitive_params: Vec<String>,
389    /// Least sensitive parameters
390    pub least_sensitive_params: Vec<String>,
391    /// Parameter interaction effects
392    pub interaction_effects: Vec<ParameterInteraction>,
393}
394
395/// Sensitivity analysis for a specific parameter
396#[derive(Debug, Clone, Serialize, Deserialize)]
397pub struct ParameterSensitivity {
398    /// Parameter name
399    pub parameter_name: String,
400    /// Current value
401    pub current_value: f64,
402    /// Sensitivity score
403    pub sensitivity_score: f64,
404    /// Optimal value range
405    pub optimal_range: (f64, f64),
406    /// Performance impact curve
407    pub impact_curve: Vec<(f64, f64)>,
408    /// Stability region
409    pub stability_region: (f64, f64),
410    /// Critical thresholds
411    pub critical_thresholds: Vec<f64>,
412}
413
414/// Interaction between parameters
415#[derive(Debug, Clone, Serialize, Deserialize)]
416pub struct ParameterInteraction {
417    /// First parameter
418    pub param1: String,
419    /// Second parameter
420    pub param2: String,
421    /// Interaction strength
422    pub interaction_strength: f64,
423    /// Interaction type
424    pub interaction_type: InteractionType,
425    /// Joint optimal region
426    pub joint_optimal_region: HashMap<String, (f64, f64)>,
427}
428
429/// Type of parameter interaction
430#[derive(Debug, Clone, Serialize, Deserialize)]
431pub enum InteractionType {
432    Synergistic,
433    Antagonistic,
434    Independent,
435    Conditional,
436}
437
438/// Architecture sensitivity analysis
439#[derive(Debug, Clone, Serialize, Deserialize)]
440pub struct ArchitectureSensitivity {
441    /// Layer depth sensitivity
442    pub depth_sensitivity: ArchitecturalSensitivity,
443    /// Layer width sensitivity
444    pub width_sensitivity: ArchitecturalSensitivity,
445    /// Attention head sensitivity
446    pub attention_head_sensitivity: ArchitecturalSensitivity,
447    /// Skip connection sensitivity
448    pub skip_connection_sensitivity: ArchitecturalSensitivity,
449    /// Architectural component importance
450    pub component_importance: HashMap<String, f64>,
451    /// Architectural bottlenecks
452    pub bottlenecks: Vec<ArchitecturalBottleneck>,
453}
454
455/// Sensitivity analysis for architectural component
456#[derive(Debug, Clone, Serialize, Deserialize)]
457pub struct ArchitecturalSensitivity {
458    /// Component name
459    pub component_name: String,
460    /// Sensitivity to changes
461    pub change_sensitivity: f64,
462    /// Performance degradation curve
463    pub degradation_curve: Vec<(f64, f64)>,
464    /// Minimum viable configuration
465    pub min_viable_config: f64,
466    /// Optimal configuration
467    pub optimal_config: f64,
468    /// Diminishing returns threshold
469    pub diminishing_returns_threshold: f64,
470}
471
472/// Architectural bottleneck
473#[derive(Debug, Clone, Serialize, Deserialize)]
474pub struct ArchitecturalBottleneck {
475    /// Bottleneck location
476    pub location: String,
477    /// Bottleneck type
478    pub bottleneck_type: BottleneckType,
479    /// Severity
480    pub severity: f64,
481    /// Performance impact
482    pub performance_impact: f64,
483    /// Resolution recommendations
484    pub resolution_recommendations: Vec<String>,
485}
486
487/// Type of architectural bottleneck
488#[derive(Debug, Clone, Serialize, Deserialize)]
489pub enum BottleneckType {
490    ComputationalBottleneck,
491    MemoryBottleneck,
492    InformationBottleneck,
493    CapacityBottleneck,
494    CommunicationBottleneck,
495}
496
497/// Data sensitivity analysis
498#[derive(Debug, Clone, Serialize, Deserialize)]
499pub struct DataSensitivity {
500    /// Training data size sensitivity
501    pub data_size_sensitivity: DataSizeSensitivity,
502    /// Data quality sensitivity
503    pub data_quality_sensitivity: DataQualitySensitivity,
504    /// Data distribution sensitivity
505    pub distribution_sensitivity: DistributionSensitivity,
506    /// Feature sensitivity analysis
507    pub feature_sensitivity: FeatureSensitivityAnalysis,
508}
509
510/// Sensitivity to training data size
511#[derive(Debug, Clone, Serialize, Deserialize)]
512pub struct DataSizeSensitivity {
513    /// Current data size
514    pub current_size: usize,
515    /// Minimum effective size
516    pub minimum_effective_size: usize,
517    /// Performance vs size curve
518    pub performance_curve: Vec<(usize, f64)>,
519    /// Data efficiency score
520    pub data_efficiency: f64,
521    /// Diminishing returns point
522    pub diminishing_returns_point: usize,
523}
524
525/// Sensitivity to data quality
526#[derive(Debug, Clone, Serialize, Deserialize)]
527pub struct DataQualitySensitivity {
528    /// Noise tolerance
529    pub noise_tolerance: f64,
530    /// Label quality importance
531    pub label_quality_importance: f64,
532    /// Feature quality importance
533    pub feature_quality_importance: f64,
534    /// Quality degradation impact
535    pub quality_impact_curve: Vec<(f64, f64)>,
536}
537
538/// Sensitivity to data distribution
539#[derive(Debug, Clone, Serialize, Deserialize)]
540pub struct DistributionSensitivity {
541    /// Distribution shift sensitivity
542    pub shift_sensitivity: f64,
543    /// Class imbalance sensitivity
544    pub imbalance_sensitivity: f64,
545    /// Domain adaptation requirements
546    pub domain_adaptation_requirements: Vec<String>,
547    /// Robustness to distribution changes
548    pub distribution_robustness: f64,
549}
550
551/// Feature-level sensitivity analysis
552#[derive(Debug, Clone, Serialize, Deserialize)]
553pub struct FeatureSensitivityAnalysis {
554    /// Most important features
555    pub most_important_features: Vec<String>,
556    /// Least important features
557    pub least_important_features: Vec<String>,
558    /// Feature interaction importance
559    pub feature_interactions: HashMap<(String, String), f64>,
560    /// Feature stability analysis
561    pub feature_stability: HashMap<String, f64>,
562}
563
564/// Training procedure sensitivity
565#[derive(Debug, Clone, Serialize, Deserialize)]
566pub struct TrainingSensitivity {
567    /// Initialization sensitivity
568    pub initialization_sensitivity: InitializationSensitivity,
569    /// Optimization method sensitivity
570    pub optimization_sensitivity: OptimizationSensitivity,
571    /// Training schedule sensitivity
572    pub schedule_sensitivity: ScheduleSensitivity,
573    /// Regularization sensitivity
574    pub regularization_sensitivity: RegularizationSensitivity,
575}
576
577/// Sensitivity to initialization
578#[derive(Debug, Clone, Serialize, Deserialize)]
579pub struct InitializationSensitivity {
580    /// Weight initialization sensitivity
581    pub weight_init_sensitivity: f64,
582    /// Bias initialization sensitivity
583    pub bias_init_sensitivity: f64,
584    /// Random seed sensitivity
585    pub seed_sensitivity: f64,
586    /// Initialization scheme importance
587    pub scheme_importance: HashMap<String, f64>,
588}
589
590/// Sensitivity to optimization method
591#[derive(Debug, Clone, Serialize, Deserialize)]
592pub struct OptimizationSensitivity {
593    /// Optimizer choice sensitivity
594    pub optimizer_sensitivity: f64,
595    /// Momentum parameter sensitivity
596    pub momentum_sensitivity: f64,
597    /// Second-order moment sensitivity
598    pub second_moment_sensitivity: f64,
599    /// Optimizer comparison
600    pub optimizer_comparison: HashMap<String, f64>,
601}
602
603/// Sensitivity to training schedule
604#[derive(Debug, Clone, Serialize, Deserialize)]
605pub struct ScheduleSensitivity {
606    /// Learning rate schedule sensitivity
607    pub lr_schedule_sensitivity: f64,
608    /// Training duration sensitivity
609    pub duration_sensitivity: f64,
610    /// Warmup sensitivity
611    pub warmup_sensitivity: f64,
612    /// Schedule parameter importance
613    pub schedule_param_importance: HashMap<String, f64>,
614}
615
616/// Sensitivity to regularization
617#[derive(Debug, Clone, Serialize, Deserialize)]
618pub struct RegularizationSensitivity {
619    /// Dropout sensitivity
620    pub dropout_sensitivity: f64,
621    /// Weight decay sensitivity
622    pub weight_decay_sensitivity: f64,
623    /// Batch normalization sensitivity
624    pub batch_norm_sensitivity: f64,
625    /// Regularization method comparison
626    pub method_comparison: HashMap<String, f64>,
627}
628
629/// Overall sensitivity insights
630#[derive(Debug, Clone, Serialize, Deserialize)]
631pub struct SensitivityInsights {
632    /// Most critical factors
633    pub most_critical_factors: Vec<String>,
634    /// Least critical factors
635    pub least_critical_factors: Vec<String>,
636    /// Surprising findings
637    pub surprising_findings: Vec<String>,
638    /// Robustness assessment
639    pub robustness_assessment: RobustnessAssessment,
640    /// Optimization recommendations
641    pub optimization_recommendations: Vec<String>,
642}
643
644/// Model robustness assessment
645#[derive(Debug, Clone, Serialize, Deserialize)]
646pub struct RobustnessAssessment {
647    /// Overall robustness score
648    pub overall_robustness: f64,
649    /// Robustness breakdown by category
650    pub category_robustness: HashMap<String, f64>,
651    /// Vulnerability areas
652    pub vulnerabilities: Vec<Vulnerability>,
653    /// Strength areas
654    pub strengths: Vec<String>,
655}
656
657/// Model vulnerability
658#[derive(Debug, Clone, Serialize, Deserialize)]
659pub struct Vulnerability {
660    /// Vulnerability type
661    pub vulnerability_type: String,
662    /// Severity level
663    pub severity: f64,
664    /// Impact description
665    pub impact: String,
666    /// Mitigation strategies
667    pub mitigation_strategies: Vec<String>,
668}
669
670/// Advanced ML debugger
671#[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    /// Create a new advanced ML debugger
680    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    /// Perform layer-wise learning rate analysis
689    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        // Analyze each layer
705        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        // Generate global insights
719        let global_lr_insights =
720            self.generate_global_lr_insights(&layer_lr_recommendations, loss_history);
721
722        // Create adaptation strategy
723        let adaptation_strategy =
724            self.create_lr_adaptation_strategy(&layer_lr_recommendations, &global_lr_insights);
725
726        // Generate training phase recommendations
727        let training_phase_recommendations =
728            self.generate_training_phase_recommendations(&adaptation_strategy);
729
730        // Predict performance with different schedules
731        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    /// Perform comprehensive model sensitivity analysis.
748    ///
749    /// **Always returns a structured error.** A sensitivity analysis measures
750    /// how performance responds when a factor is *varied*, which requires
751    /// re-evaluating the model at perturbed hyperparameters, architectures,
752    /// dataset sizes and seeds. This method receives only the CURRENT parameter
753    /// values plus a flat `&[f64]` of already-observed metrics -- one point per
754    /// factor -- so no response curve, optimal range, stability region or
755    /// robustness score is derivable from its inputs.
756    ///
757    /// It used to return a fully populated [`ModelSensitivityAnalysisResult`]
758    /// assembled from hardcoded constants: `sensitivity_score: 0.8`,
759    /// `optimal_range: (0.0001, 0.01)`, `current_size: 10000`,
760    /// `most_important_features: ["feature_1", "feature_2"]`, and a
761    /// `surprising_findings` string, none of which looked at `model_params`,
762    /// `performance_metrics` or `architecture_config`. The five helpers that
763    /// produced those constants have been deleted.
764    ///
765    /// The result types remain public so a caller that really does run a
766    /// parameter sweep can construct and share one.
767    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    /// Generate comprehensive advanced ML debugging report
788    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    // Helper methods for layer-wise LR analysis
807
808    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        // Calculate gradient statistics
817        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        // Estimate optimal learning rate based on gradient properties
823        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            // Adaptive learning rate based on gradient properties
830            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        // Calculate layer metrics
839        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        // Determine urgency
850        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        // Generate reasoning
862        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            // No statistical confidence is derivable here: the recommendation
876            // comes from one gradient/weight snapshot of one layer, with no
877            // repeated trials and no held-out evaluation to estimate a
878            // sampling distribution from. Previously a hardcoded `0.8`.
879            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        // Simplified strategy creation
921        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    // Helper methods for sensitivity analysis
995
996    // Additional helper methods
997
998    fn estimate_layer_loss_contribution(&self, loss_history: &[f64]) -> f64 {
999        // Simplified estimation
1000        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) // Higher stability for lower variance
1019    }
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) // Better health for lower variance
1061    }
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, // Slightly pessimistic prediction
1090        }
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/// Comprehensive advanced ML debugging report
1149#[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    // ---- Wave 6c debug-sweep2 honesty regressions ------------------------
1165
1166    #[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(&params, &[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        // Create test data
1239        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}