ujjwalpardeshi/pytorch-training-debugger
2
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 