Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
environment.py536 linesDownload Raw Back to server
1"""MLTrainingEnvironment โ€” extends openenv Environment.2 3Session isolation, progressive information reveal, error handling.4step() never raises an unhandled exception.5"""6 7from __future__ import annotations8 9import dataclasses10import logging11import uuid12from typing import Any, Optional, Union13 14import torch15from openenv.core.env_server.interfaces import Environment16 17from ml_training_debugger.code_templates import (18    generate_code_snippet,19    validate_fix,20)21from ml_training_debugger.graders import grade_episode22from ml_training_debugger.models import (23    ALL_ACTION_TYPES,24    VALID_CONFIG_KEYS,25    VALID_DIAGNOSES,26    CodeSnippet,27    DataBatchStats,28    EpisodeState,29    GradientStats,30    MLTrainingAction,31    MLTrainingObservation,32    ModelWeightStats,33    TrainingConfig,34)35from ml_training_debugger.pytorch_engine import (36    create_model_and_inject_fault,37    extract_gradient_stats,38    extract_model_modes,39    extract_weight_stats,40)41from ml_training_debugger.reward_engine import compute_reward42from ml_training_debugger.scenarios import ScenarioParams, sample_scenario43from ml_training_debugger.simulation import (44    gen_data_batch_stats,45    gen_loss_history,46    gen_val_accuracy_history,47    gen_val_loss_history,48)49from server._baseline_results import store_grader_result50 51logger = logging.getLogger(__name__)52 53 54@dataclasses.dataclass55class SessionData:56    """Per-session episode data."""57 58    scenario: ScenarioParams59    model: torch.nn.Module60    state: EpisodeState61    config: TrainingConfig62    gradient_stats: list[GradientStats]63    weight_stats: list[ModelWeightStats] | None64    model_modes: dict[str, str] | None65    data_batch_stats_raw: dict[str, Union[int, float, list, dict, None]] | None66    code_snippet_raw: dict[str, Union[str, int, list, None]] | None67    loss_history: list[float]68    val_acc_history: list[float]69    val_loss_history: list[float]70    done: bool71    last_score: float | None72    convergence_after_fix: bool73 74 75class MLTrainingEnvironment(Environment[MLTrainingAction, MLTrainingObservation, dict]):76    """OpenEnv environment for PyTorch training run debugging."""77 78    SUPPORTS_CONCURRENT_SESSIONS = True79 80    def __init__(self, **kwargs: Any) -> None:81        super().__init__(**kwargs)82        self._sessions: dict[str, SessionData] = {}83        self._last_completed: dict[str, dict] = {}84        self._current_session_id: str = ""85 86    def _get_session(self, episode_id: str | None = None) -> SessionData | None:87        sid = episode_id or self._current_session_id88        return self._sessions.get(sid)89 90    def _build_observation(91        self, session: SessionData, reward: float = 0.092    ) -> MLTrainingObservation:93        """Build observation from session data."""94        state = session.state95 96        gradient_stats_models = []97        if state.gradients_inspected and session.gradient_stats:98            gradient_stats_models = session.gradient_stats99 100        weight_stats_models = None101        if state.model_weights_inspected and session.weight_stats is not None:102            weight_stats_models = session.weight_stats103 104        data_batch = None105        if state.data_inspected and session.data_batch_stats_raw is not None:106            data_batch = DataBatchStats(**session.data_batch_stats_raw)107 108        model_modes = None109        if state.model_modes_inspected and session.model_modes is not None:110            model_modes = session.model_modes111 112        code_snippet = None113        if state.code_inspected and session.code_snippet_raw is not None:114            code_snippet = CodeSnippet(**session.code_snippet_raw)115 116        return MLTrainingObservation(117            run_id=self._current_session_id,118            framework="pytorch",119            epoch=20,120            training_loss_history=session.loss_history,121            val_loss_history=session.val_loss_history,122            val_accuracy_history=session.val_acc_history,123            gradient_stats=gradient_stats_models,124            model_weight_stats=weight_stats_models,125            gpu_memory_used_gb=session.scenario.gpu_memory_used_gb,126            gpu_memory_total_gb=16.0,127            learning_rate=session.config.learning_rate,128            current_config=session.config,129            error_log=session.scenario.error_log,130            data_batch_stats=data_batch,131            model_mode_info=model_modes,132            code_snippet=code_snippet,133            available_actions=state.compute_available_actions(),134            episode_state=state,135            notes=session.scenario.notes,136            done=session.done,137            reward=reward,138        )139 140    def reset(141        self,142        seed: Optional[int] = None,143        episode_id: Optional[str] = None,144        **kwargs: Any,145    ) -> MLTrainingObservation:146        """Reset environment for a new episode."""147        # Determine task_id โ€” passed via kwargs or defaults to task_001148        task_id = kwargs.get("task_id", "task_001")149 150        # If called with episode_id that has an active session, terminate it151        session_id = episode_id or str(uuid.uuid4())152        if session_id in self._sessions:153            old = self._sessions[session_id]154            if not old.done:155                score = grade_episode(old.scenario.task_id, old.state, old.scenario)156                self._last_completed[session_id] = {157                    "score": score,158                    "task_id": old.scenario.task_id,159                    "steps": old.state.step_count,160                }161                store_grader_result(162                    session_id, score, old.scenario.task_id, old.state.step_count163                )164 165        self._current_session_id = session_id166 167        # Derive deterministic seed and difficulty168        base_seed = seed if seed is not None else 42169        difficulty_level = kwargs.get("difficulty_level", 3)170        scenario = sample_scenario(task_id, base_seed, difficulty_level=difficulty_level)171 172        # Set torch seed for reproducibility173        torch.manual_seed(scenario.seed)174 175        # Create real PyTorch model with fault injection176        model, info = create_model_and_inject_fault(scenario)177 178        # Generate parametric curves179        loss_history = gen_loss_history(scenario)180        val_acc_history = gen_val_accuracy_history(scenario)181        val_loss_history = gen_val_loss_history(scenario)182 183        # Pre-generate data batch stats184        data_batch_raw = gen_data_batch_stats(scenario)185 186        # Pre-generate code snippet (for Task 6)187        code_snippet_raw = None188        if scenario.bug_type is not None:189            code_snippet_raw = generate_code_snippet(scenario.bug_type, scenario.seed)190 191        # Build initial config from scenario192        config = TrainingConfig(193            learning_rate=scenario.learning_rate,194            weight_decay=scenario.weight_decay,195        )196 197        # Create fresh episode state198        state = EpisodeState()199 200        session = SessionData(201            scenario=scenario,202            model=model,203            state=state,204            config=config,205            gradient_stats=[],206            weight_stats=None,207            model_modes=None,208            data_batch_stats_raw=data_batch_raw,209            code_snippet_raw=code_snippet_raw,210            loss_history=loss_history,211            val_acc_history=val_acc_history,212            val_loss_history=val_loss_history,213            done=False,214            last_score=None,215            convergence_after_fix=False,216        )217 218        self._sessions[session_id] = session219 220        logger.info(221            "reset",222            extra={223                "session_id": session_id,224                "task_id": task_id,225                "scenario_seed": scenario.seed,226            },227        )228 229        return self._build_observation(session)230 231    def step(232        self,233        action: MLTrainingAction,234        timeout_s: Optional[float] = None,235        **kwargs: Any,236    ) -> MLTrainingObservation:237        """Process one agent action. Never raises."""238        session = self._get_session()239 240        # No active episode241        if session is None:242            return MLTrainingObservation(243                done=True,244                reward=0.0,245                error_log="Error: no active episode. Call reset(task_id) first.",246            )247 248        # Episode already done249        if session.done:250            return self._build_observation(session, reward=0.0)251 252        state = session.state253        scenario = session.scenario254        action_type = action.action_type255 256        # Increment step count257        state.step_count += 1258 259        # Validate action_type is a known type260        if action_type not in ALL_ACTION_TYPES:261            reward = compute_reward(action, state, scenario, is_valid_action=False)262            state.actions_taken.append(f"invalid:{action_type}")263            obs = self._build_observation(session, reward=reward)264            obs.error_log = (265                f"Invalid action_type: {action_type}. "266                f"Valid types: {sorted(ALL_ACTION_TYPES)}"267            )268            return obs269 270        # Check if action is in available_actions271        available = state.compute_available_actions()272        if action_type not in available:273            reward = compute_reward(action, state, scenario, is_valid_action=False)274            state.actions_taken.append(f"unavailable:{action_type}")275            obs = self._build_observation(session, reward=reward)276            obs.error_log = (277                f"Action '{action_type}' not available. " f"Available: {available}"278            )279            return obs280 281        # Validate required fields for specific actions282        error = self._validate_action_fields(action)283        if error is not None:284            reward = compute_reward(action, state, scenario, is_valid_action=False)285            state.actions_taken.append(f"malformed:{action_type}")286            obs = self._build_observation(session, reward=reward)287            obs.error_log = error288            return obs289 290        # Dispatch action291        is_correct_fix: bool | None = None292        convergence = False293 294        # Snapshot state BEFORE dispatch โ€” reward engine needs pre-action state295        # to correctly compute investigation bonuses and context-gated penalties296        state_before = state.model_copy(deep=True)297 298        try:299            is_correct_fix, convergence = self._dispatch_action(action, session)300        except Exception as exc:301            logger.error(302                "step_error",303                extra={304                    "session_id": self._current_session_id,305                    "action": action_type,306                    "error": str(exc),307                },308                exc_info=True,309            )310            reward = compute_reward(action, state_before, scenario, is_valid_action=False)311            obs = self._build_observation(session, reward=reward)312            obs.error_log = f"Internal error processing {action_type}: {exc}"313            return obs314 315        # Record action316        if action_type == "mark_diagnosed" and action.diagnosis:317            state.actions_taken.append(f"mark_diagnosed:{action.diagnosis}")318        else:319            state.actions_taken.append(action_type)320 321        # Compute reward using pre-action state322        reward = compute_reward(323            action,324            state_before,325            scenario,326            is_valid_action=True,327            is_correct_fix=is_correct_fix,328            convergence_confirmed=convergence,329        )330 331        # Check step limit332        if state.step_count >= scenario.max_steps and not session.done:333            session.done = True334 335        # Check done336        if session.done:337            score = grade_episode(scenario.task_id, state, scenario)338            session.last_score = score339            self._last_completed[self._current_session_id] = {340                "score": score,341                "task_id": scenario.task_id,342                "steps": state.step_count,343            }344            store_grader_result(345                self._current_session_id, score, scenario.task_id, state.step_count346            )347            logger.info(348                "episode_completed",349                extra={350                    "session_id": self._current_session_id,351                    "task_id": scenario.task_id,352                    "steps": state.step_count,353                    "score": score,354                },355            )356 357        logger.info(358            "step",359            extra={360                "session_id": self._current_session_id,361                "step_count": state.step_count,362                "action_type": action_type,363                "reward": reward,364            },365        )366 367        return self._build_observation(session, reward=reward)368 369    def _validate_action_fields(self, action: MLTrainingAction) -> str | None:370        """Validate required fields for specific actions. Return error or None."""371        if action.action_type == "modify_config":372            if action.target is None or action.value is None:373                return "modify_config requires 'target' and 'value' fields"374            if action.target not in VALID_CONFIG_KEYS:375                return f"Unknown config key: {action.target}. Valid: {sorted(VALID_CONFIG_KEYS)}"376 377        if action.action_type == "mark_diagnosed":378            if action.diagnosis is None:379                return "mark_diagnosed requires 'diagnosis' field"380            if action.diagnosis not in VALID_DIAGNOSES:381                return (382                    f"Invalid diagnosis: {action.diagnosis}. "383                    f"Valid: {sorted(VALID_DIAGNOSES)}"384                )385 386        if action.action_type == "fix_code":387            if action.line is None or action.replacement is None:388                return "fix_code requires 'line' and 'replacement' fields"389 390        return None391 392    def _dispatch_action(393        self, action: MLTrainingAction, session: SessionData394    ) -> tuple[bool | None, bool]:395        """Dispatch action to handler. Returns (is_correct_fix, convergence)."""396        state = session.state397        scenario = session.scenario398        is_correct_fix: bool | None = None399        convergence = False400 401        at = action.action_type402 403        if at == "inspect_gradients":404            if not state.gradients_inspected:405                stats = extract_gradient_stats(session.model, scenario)406                session.gradient_stats = stats407                state.gradients_inspected = True408                # Set gradients_were_normal: True if ALL layers is_exploding=False409                state.gradients_were_normal = all(not s.is_exploding for s in stats)410 411        elif at == "inspect_data_batch":412            state.data_inspected = True413 414        elif at == "inspect_model_modes":415            if not state.model_modes_inspected:416                modes = extract_model_modes(session.model)417                session.model_modes = modes418                state.model_modes_inspected = True419 420        elif at == "inspect_model_weights":421            if not state.model_weights_inspected:422                stats = extract_weight_stats(session.model)423                session.weight_stats = stats424                state.model_weights_inspected = True425 426        elif at == "inspect_code":427            state.code_inspected = True428 429        elif at == "modify_config":430            if action.target and action.value is not None:431                setattr(session.config, action.target, action.value)432                state.fix_action_taken = True433 434        elif at == "add_callback":435            state.fix_action_taken = True436 437        elif at == "replace_optimizer":438            state.fix_action_taken = True439 440        elif at == "patch_data_loader":441            state.fix_action_taken = True442 443        elif at == "fix_model_mode":444            state.fix_action_taken = True445 446        elif at == "fix_code":447            state.fix_action_taken = True448            if scenario.bug_type and action.line and action.replacement:449                is_correct_fix = validate_fix(450                    scenario.bug_type, action.line, action.replacement451                )452            else:453                is_correct_fix = False454 455        elif at == "restart_run":456            state.restart_after_fix = True457            # Check convergence โ€” did the fix address the root cause?458            convergence = self._check_convergence(session)459            session.convergence_after_fix = convergence460 461        elif at == "mark_diagnosed":462            state.diagnosis_submitted = True463            session.done = True464 465        return is_correct_fix, convergence466 467    def _check_convergence(self, session: SessionData) -> bool:468        """Check if the applied fix would resolve the root cause."""469        scenario = session.scenario470        state = session.state471        root = scenario.root_cause.value472 473        if root == "lr_too_high":474            return (475                "modify_config" in state.actions_taken476                and session.config.learning_rate <= 0.001477            )478 479        if root == "vanishing_gradients":480            return (481                "modify_config" in state.actions_taken482                and session.config.learning_rate >= 0.001483            )484 485        if root == "data_leakage":486            return "patch_data_loader" in state.actions_taken487 488        if root == "overfitting":489            return (490                "modify_config" in state.actions_taken491                or "add_callback" in state.actions_taken492            )493 494        if root == "batchnorm_eval_mode":495            return "fix_model_mode" in state.actions_taken496 497        if root == "code_bug":498            return "fix_code" in state.actions_taken and state.fix_action_taken499 500        if root == "scheduler_misconfigured":501            return "modify_config" in state.actions_taken502 503        return False504 505    @property506    def state(self) -> dict:507        """Return current environment state."""508        session = self._get_session()509        if session is None:510            return {"status": "no_active_episode"}511        st = session.state512        return {513            "status": "active",514            "task_id": session.scenario.task_id,515            "step_count": st.step_count,516            "done": session.done,517            "gradients_inspected": st.gradients_inspected,518            "data_inspected": st.data_inspected,519            "model_modes_inspected": st.model_modes_inspected,520            "model_weights_inspected": st.model_weights_inspected,521            "code_inspected": st.code_inspected,522            "fix_action_taken": st.fix_action_taken,523            "restart_after_fix": st.restart_after_fix,524            "diagnosis_submitted": st.diagnosis_submitted,525            "available_actions": st.compute_available_actions(),526        }527 528    def get_last_completed(self, session_id: str | None = None) -> dict | None:529        """Get last completed episode data for grader endpoint."""530        if session_id:531            return self._last_completed.get(session_id)532        # Return most recent533        if self._last_completed:534            return list(self._last_completed.values())[-1]535        return None536