ujjwalpardeshi/pytorch-training-debugger
2
1"""All Pydantic models, enums, and typed data structures.2 3No business logic. Pure data definitions.4"""5 6from __future__ import annotations7 8import enum9from typing import Optional, Union10 11import torch # noqa: F40112from openenv.core.env_server.types import Action, Observation13from pydantic import BaseModel, Field14 15 16class RootCauseDiagnosis(str, enum.Enum):17 """Closed enumeration of ML failure root causes."""18 19 LR_TOO_HIGH = "lr_too_high"20 VANISHING_GRADIENTS = "vanishing_gradients"21 DATA_LEAKAGE = "data_leakage"22 OVERFITTING = "overfitting"23 BATCHNORM_EVAL_MODE = "batchnorm_eval_mode"24 CODE_BUG = "code_bug"25 SCHEDULER_MISCONFIGURED = "scheduler_misconfigured"26 27 28VALID_DIAGNOSES: set[str] = {d.value for d in RootCauseDiagnosis}29 30 31class TrainingConfig(BaseModel):32 """Typed hyperparameter configuration."""33 34 learning_rate: float = 0.00135 weight_decay: float = 0.000136 batch_size: int = 6437 hidden_dim: int = 6438 num_layers: int = 339 optimizer: str = "adam"40 dropout_rate: float = 0.041 gradient_clip_norm: Optional[float] = None42 43 44VALID_CONFIG_KEYS: set[str] = set(TrainingConfig.model_fields.keys())45 46 47class GradientStats(BaseModel):48 """Per-layer gradient information from real torch.autograd."""49 50 layer_name: str51 norm_history: list[float]52 mean_norm: float53 max_norm: float54 is_exploding: bool # True when mean_norm > 10.055 is_vanishing: bool # True when mean_norm < 1e-656 57 58class ModelWeightStats(BaseModel):59 """Per-layer weight statistics from real state_dict()."""60 61 layer_name: str62 weight_norm: float63 weight_mean: float64 weight_std: float65 weight_min: float66 weight_max: float67 dead_neuron_pct: float = 0.068 has_nan: bool = False69 has_inf: bool = False70 71 72class DataBatchStats(BaseModel):73 """Data batch inspection results."""74 75 label_distribution: dict[int, float]76 feature_mean: float77 feature_std: float78 null_count: int = 079 class_overlap_score: float80 batch_size: int81 duplicate_ratio: float = 0.082 confusion_matrix: Optional[list[list[float]]] = None83 84 85class CodeSnippet(BaseModel):86 """PyTorch code for Task 6 inspection."""87 88 code: str89 filename: str = "train.py"90 line_count: int91 imports: list[str]92 hint: Optional[str] = None93 94 95class EpisodeState(BaseModel):96 """Tracks agent history within an episode."""97 98 step_count: int = 099 gradients_inspected: bool = False100 gradients_were_normal: bool = False101 data_inspected: bool = False102 model_modes_inspected: bool = False103 model_weights_inspected: bool = False104 code_inspected: bool = False105 fix_action_taken: bool = False106 restart_after_fix: bool = False107 diagnosis_submitted: bool = False108 actions_taken: list[str] = Field(default_factory=list)109 110 def compute_available_actions(self) -> list[str]:111 """Dynamically compute available actions based on current state."""112 actions: list[str] = [113 "inspect_gradients",114 "inspect_data_batch",115 "inspect_model_modes",116 "inspect_model_weights",117 "inspect_code",118 "modify_config",119 "add_callback",120 "replace_optimizer",121 "patch_data_loader",122 "fix_model_mode",123 ]124 if self.code_inspected:125 actions.append("fix_code")126 if self.fix_action_taken:127 actions.append("restart_run")128 if not self.diagnosis_submitted:129 actions.append("mark_diagnosed")130 return actions131 132 133ALL_ACTION_TYPES: set[str] = {134 "inspect_gradients",135 "inspect_data_batch",136 "inspect_model_modes",137 "inspect_model_weights",138 "inspect_code",139 "modify_config",140 "add_callback",141 "replace_optimizer",142 "patch_data_loader",143 "fix_model_mode",144 "fix_code",145 "restart_run",146 "mark_diagnosed",147}148 149 150class MLTrainingAction(Action):151 """What the agent can do — extends openenv Action."""152 153 action_type: str154 target: Optional[str] = None155 value: Optional[Union[float, int, str]] = None156 diagnosis: Optional[str] = None157 line: Optional[int] = None158 replacement: Optional[str] = None159 160 161class MLTrainingObservation(Observation):162 """Full observation — extends openenv Observation.163 164 Observation base has built-in: done (bool), reward (float|None), metadata (dict).165 """166 167 run_id: str = ""168 framework: str = "pytorch"169 epoch: int = 20170 training_loss_history: list[float] = Field(default_factory=list)171 val_loss_history: list[float] = Field(default_factory=list)172 val_accuracy_history: list[float] = Field(default_factory=list)173 gradient_stats: list[GradientStats] = Field(default_factory=list)174 model_weight_stats: Optional[list[ModelWeightStats]] = None175 gpu_memory_used_gb: float = 6.2176 gpu_memory_total_gb: float = 16.0177 learning_rate: float = 0.001178 current_config: TrainingConfig = Field(default_factory=TrainingConfig)179 error_log: Optional[str] = None180 data_batch_stats: Optional[DataBatchStats] = None181 model_mode_info: Optional[dict[str, str]] = None182 code_snippet: Optional[CodeSnippet] = None183 available_actions: list[str] = Field(default_factory=list)184 episode_state: EpisodeState = Field(default_factory=EpisodeState)185 notes: Optional[str] = None186 