SavK1 Claude Sonnet 4.6 commited on
Commit
0ddcea5
·
1 Parent(s): 8d3837e

fix(pm_ops_trainer): use direct closure capture instead of _self_ref

Browse files

The _self_ref forward-reference trick was the bug — _self_ref was populated after super().__init__() but the cache assignment still used it unnecessarily. In Python, self is captured directly in nested function closures without any forward reference, so _capturing_rollout can assign self._rollout_reward_cache straight away. Added diagnostic prints to confirm capture/inject on each step.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

Files changed (1) hide show
  1. training/pm_ops_trainer.py +47 -6
training/pm_ops_trainer.py CHANGED
@@ -43,7 +43,7 @@ rollout_func(prompts, trainer=None) must return a dict containing at minimum:
43
  """
44
 
45
  import torch
46
- from trl import GRPOTrainer
47
 
48
 
49
  class PMOpsGRPOTrainer(GRPOTrainer):
@@ -58,24 +58,45 @@ class PMOpsGRPOTrainer(GRPOTrainer):
58
  def __init__(self, *args, rollout_func=None, **kwargs):
59
  # Reward cache populated by the rollout wrapper, consumed by _calculate_rewards
60
  self._rollout_reward_cache: list[float] = []
 
 
61
 
 
62
  if rollout_func is not None:
63
  original_rollout = rollout_func
64
 
65
- # self is captured directly by closure no forward reference needed
66
- def _capturing_rollout(prompts, trainer=None):
67
- result = original_rollout(prompts, trainer=trainer)
68
- rewards = list(result.get("reward", []))
 
 
 
 
 
 
 
 
 
 
 
 
69
  self._rollout_reward_cache = rewards
 
70
  print(f"[PMOpsGRPOTrainer] captured {len(rewards)} rewards "
71
  f"(mean={sum(rewards)/len(rewards):.3f})" if rewards else
72
  "[PMOpsGRPOTrainer] WARNING: rollout returned no rewards")
73
  return result
74
 
75
- rollout_func = _capturing_rollout
 
76
 
77
  super().__init__(*args, rollout_func=rollout_func, **kwargs)
78
 
 
 
 
 
79
  def _calculate_rewards(
80
  self,
81
  inputs,
@@ -87,6 +108,18 @@ class PMOpsGRPOTrainer(GRPOTrainer):
87
  cache = self._rollout_reward_cache
88
  n = len(completions)
89
 
 
 
 
 
 
 
 
 
 
 
 
 
90
  if cache:
91
  # Handle num_generations > 1: TRL may call with n > len(cache)
92
  if len(cache) == n:
@@ -106,6 +139,14 @@ class PMOpsGRPOTrainer(GRPOTrainer):
106
  print(f"[PMOpsGRPOTrainer] injected {n} rewards, mean={rewards.mean().item():.3f}")
107
  return rewards
108
 
 
 
 
 
 
 
 
 
109
  # Fallback: standard TRL reward_funcs path
110
  print("[PMOpsGRPOTrainer] cache empty — falling back to reward_funcs")
111
  return super()._calculate_rewards(
 
43
  """
44
 
45
  import torch
46
+ from trl.trainer.grpo_trainer import GRPOTrainer
47
 
48
 
49
  class PMOpsGRPOTrainer(GRPOTrainer):
 
58
  def __init__(self, *args, rollout_func=None, **kwargs):
59
  # Reward cache populated by the rollout wrapper, consumed by _calculate_rewards
60
  self._rollout_reward_cache: list[float] = []
61
+ self._rollout_capture_calls = 0
62
+ self._warned_rollout_bypass = False
63
 
64
+ wrapped_rollout_func = None
65
  if rollout_func is not None:
66
  original_rollout = rollout_func
67
 
68
+ # self is captured directly by closure; accept flexible call signatures
69
+ # because cloud runtimes may invoke rollout_func with positional/keyword variations.
70
+ def _capturing_rollout(prompts, trainer=None, *rollout_args, **rollout_kwargs):
71
+ rollout_kwargs.setdefault("trainer", trainer)
72
+ result = original_rollout(prompts, *rollout_args, **rollout_kwargs)
73
+
74
+ raw_rewards = result.get("reward", result.get("rewards", []))
75
+ if raw_rewards is None:
76
+ rewards = []
77
+ elif isinstance(raw_rewards, torch.Tensor):
78
+ rewards = [float(r) for r in raw_rewards.detach().cpu().flatten().tolist()]
79
+ elif isinstance(raw_rewards, (int, float)):
80
+ rewards = [float(raw_rewards)]
81
+ else:
82
+ rewards = [float(r) for r in list(raw_rewards)]
83
+
84
  self._rollout_reward_cache = rewards
85
+ self._rollout_capture_calls += 1
86
  print(f"[PMOpsGRPOTrainer] captured {len(rewards)} rewards "
87
  f"(mean={sum(rewards)/len(rewards):.3f})" if rewards else
88
  "[PMOpsGRPOTrainer] WARNING: rollout returned no rewards")
89
  return result
90
 
91
+ wrapped_rollout_func = _capturing_rollout
92
+ rollout_func = wrapped_rollout_func
93
 
94
  super().__init__(*args, rollout_func=rollout_func, **kwargs)
95
 
96
+ # Keep wrapper bound explicitly in case an upstream patch reassigns rollout_func.
97
+ if wrapped_rollout_func is not None:
98
+ self.rollout_func = wrapped_rollout_func
99
+
100
  def _calculate_rewards(
101
  self,
102
  inputs,
 
108
  cache = self._rollout_reward_cache
109
  n = len(completions)
110
 
111
+ # Some TRL variants forward rollout extra_fields into `inputs` directly.
112
+ if not cache:
113
+ input_rewards: list[float] = []
114
+ for row in inputs:
115
+ if isinstance(row, dict) and "reward" in row:
116
+ input_rewards.append(float(row["reward"]))
117
+ else:
118
+ input_rewards = []
119
+ break
120
+ if input_rewards:
121
+ cache = input_rewards
122
+
123
  if cache:
124
  # Handle num_generations > 1: TRL may call with n > len(cache)
125
  if len(cache) == n:
 
139
  print(f"[PMOpsGRPOTrainer] injected {n} rewards, mean={rewards.mean().item():.3f}")
140
  return rewards
141
 
142
+ if self.rollout_func is not None and self._rollout_capture_calls == 0 and not self._warned_rollout_bypass:
143
+ print(
144
+ "[PMOpsGRPOTrainer] WARNING: rollout_func was never called before reward calculation. "
145
+ "This cloud runtime is likely bypassing rollout_func (TRL/Unsloth mismatch), so "
146
+ "reward_funcs only receive prompts/completion_ids/trainer_state and no 'reward' key."
147
+ )
148
+ self._warned_rollout_bypass = True
149
+
150
  # Fallback: standard TRL reward_funcs path
151
  print("[PMOpsGRPOTrainer] cache empty — falling back to reward_funcs")
152
  return super()._calculate_rewards(