ujjwalpardeshi/pytorch-training-debugger
2
1"""Reward function — 7 components, separate from graders.2 3Returns a float per step for RL training signal. Hard cap at [-1.0, 1.0].4"""5 6from __future__ import annotations7 8import torch # noqa: F4019 10from ml_training_debugger.models import EpisodeState, MLTrainingAction11from ml_training_debugger.scenarios import ScenarioParams12 13# Reward constants14STEP_PENALTY = -0.0115INVESTIGATION_BONUS = 0.0516CONTEXT_GATED_PENALTY = -0.2017INVALID_ACTION_PENALTY = -0.0518WRONG_CODE_FIX_PENALTY = -0.1019CORRECT_DIAGNOSIS_REWARD = 0.5020WRONG_DIAGNOSIS_PENALTY = -0.3021TERMINAL_CONVERGENCE_REWARD = 0.4022 23INVESTIGATION_ACTIONS = frozenset(24 {25 "inspect_gradients",26 "inspect_data_batch",27 "inspect_model_modes",28 "inspect_model_weights",29 "inspect_code",30 }31)32 33_INSPECTION_STATE_MAP = {34 "inspect_gradients": "gradients_inspected",35 "inspect_data_batch": "data_inspected",36 "inspect_model_modes": "model_modes_inspected",37 "inspect_model_weights": "model_weights_inspected",38 "inspect_code": "code_inspected",39}40 41 42def compute_reward(43 action: MLTrainingAction,44 state: EpisodeState,45 scenario: ScenarioParams,46 is_valid_action: bool = True,47 is_correct_fix: bool | None = None,48 convergence_confirmed: bool = False,49) -> float:50 """Compute reward for a single step.51 52 Args:53 action: The action taken.54 state: Episode state BEFORE the action is applied.55 scenario: Current scenario params.56 is_valid_action: Whether the action is in available_actions.57 is_correct_fix: For fix_code — True/False/None.58 convergence_confirmed: Whether restart showed convergence.59 60 Returns:61 Reward float, capped at [-1.0, 1.0].62 """63 reward = 0.064 65 # Step penalty (unconditional)66 reward += STEP_PENALTY67 68 if not is_valid_action:69 reward += INVALID_ACTION_PENALTY70 return max(-1.0, min(1.0, reward))71 72 action_type = action.action_type73 74 # Investigation bonus — first-time only75 if action_type in INVESTIGATION_ACTIONS:76 state_field = _INSPECTION_STATE_MAP.get(action_type)77 if state_field and not getattr(state, state_field):78 reward += INVESTIGATION_BONUS79 80 # Context-gated penalty: adding gradient clipping after seeing normal gradients81 if action_type == "add_callback":82 if state.gradients_inspected and state.gradients_were_normal:83 reward += CONTEXT_GATED_PENALTY84 85 # Wrong code fix86 if action_type == "fix_code" and is_correct_fix is False:87 reward += WRONG_CODE_FIX_PENALTY88 89 # Diagnosis90 if action_type == "mark_diagnosed":91 if action.diagnosis == scenario.root_cause.value:92 reward += CORRECT_DIAGNOSIS_REWARD93 else:94 reward += WRONG_DIAGNOSIS_PENALTY95 96 # Convergence after fix+restart97 if action_type == "restart_run":98 if state.fix_action_taken and convergence_confirmed:99 reward += TERMINAL_CONVERGENCE_REWARD100 101 return max(-1.0, min(1.0, reward))102 