Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
reward_engine.py102 linesDownload Raw Back to ml_training_debugger
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