Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
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