Spaces:
Sleeping
fix(grpo): eliminate zero-advantage collapse from stale rewards
Browse filesThree root causes all produced reward=const β advantage=0 β zero gradient:
1. generate_rollout_completions uses near-greedy decoding, making every
rollout for the same prompt produce identical outputs. Replaced with
_generate_no_vllm at temperature=1.1 (top_p=0.95, top_k=50).
2. Reward formula env_score*0.85 + json_ratio*0.15 collapses to 0.15
whenever env_score=0 and json_ratio=1.0 (always true after SFT). Added
ok_ratio (fraction of steps env accepted) as the variance source when
the task isn't yet solved; ok_ratio differs across rollouts even with
identical JSON output.
3. All num_generations copies of the same prompt hit the same env seed β
same starting state β same reward. Added gen_offset (tracked per prompt
in grpo_rollout_func) so each generation resets the env with seed+offset,
exploring a distinct org config and task variant.
Also fixed _calculate_rewards tile-mod (cache[i % k] gave every generation
of the same prompt an identical reward even if rollouts differed). Now pads
with mean instead of repeating, and logs reward/injected_std to confirm
variance is non-zero.
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
- training/pm_ops_trainer.py +18 -8
- training/rollout.py +40 -15
- training/train_v3.ipynb +2 -101
|
@@ -130,12 +130,13 @@ class PMOpsGRPOTrainer(GRPOTrainer):
|
|
| 130 |
cache = input_rewards
|
| 131 |
|
| 132 |
# Cloud fallback: some patched runtimes skip the normal rollout capture path
|
| 133 |
-
# before calling _calculate_rewards. Actively invoke rollout_func
|
| 134 |
-
#
|
|
|
|
| 135 |
if not cache and self.rollout_func is not None and prompts:
|
| 136 |
try:
|
| 137 |
-
print("[PMOpsGRPOTrainer] cache empty β probing rollout_func for rewards")
|
| 138 |
-
out = self.rollout_func(prompts, trainer=self)
|
| 139 |
cache = self._rollout_reward_cache
|
| 140 |
if not cache and isinstance(out, dict):
|
| 141 |
cache = _coerce_rewards(out.get("reward", out.get("rewards", [])))
|
|
@@ -144,11 +145,17 @@ class PMOpsGRPOTrainer(GRPOTrainer):
|
|
| 144 |
print(f"[PMOpsGRPOTrainer] rollout probe failed: {exc!r}")
|
| 145 |
|
| 146 |
if cache:
|
| 147 |
-
# Handle num_generations > 1: TRL may call with n > len(cache)
|
| 148 |
if len(cache) == n:
|
| 149 |
rewards_list = cache
|
|
|
|
|
|
|
| 150 |
else:
|
| 151 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
|
| 153 |
device = self.accelerator.device
|
| 154 |
rewards = torch.tensor(
|
|
@@ -157,9 +164,12 @@ class PMOpsGRPOTrainer(GRPOTrainer):
|
|
| 157 |
device=device,
|
| 158 |
).unsqueeze(1) # [batch_size, 1]
|
| 159 |
|
| 160 |
-
|
|
|
|
|
|
|
| 161 |
self._rollout_reward_cache = [] # consume cache
|
| 162 |
-
print(f"[PMOpsGRPOTrainer] injected {n} rewards
|
|
|
|
| 163 |
return rewards
|
| 164 |
|
| 165 |
if self.rollout_func is not None and self._rollout_capture_calls == 0 and not self._warned_rollout_bypass:
|
|
|
|
| 130 |
cache = input_rewards
|
| 131 |
|
| 132 |
# Cloud fallback: some patched runtimes skip the normal rollout capture path
|
| 133 |
+
# before calling _calculate_rewards. Actively invoke rollout_func here.
|
| 134 |
+
# Pass ALL n prompts (including repeated ones for num_generations > 1) so
|
| 135 |
+
# the rollout_func can generate distinct rewards per generation β NOT tile-mod.
|
| 136 |
if not cache and self.rollout_func is not None and prompts:
|
| 137 |
try:
|
| 138 |
+
print(f"[PMOpsGRPOTrainer] cache empty β probing rollout_func for {n} rewards")
|
| 139 |
+
out = self.rollout_func(list(prompts), trainer=self)
|
| 140 |
cache = self._rollout_reward_cache
|
| 141 |
if not cache and isinstance(out, dict):
|
| 142 |
cache = _coerce_rewards(out.get("reward", out.get("rewards", [])))
|
|
|
|
| 145 |
print(f"[PMOpsGRPOTrainer] rollout probe failed: {exc!r}")
|
| 146 |
|
| 147 |
if cache:
|
|
|
|
| 148 |
if len(cache) == n:
|
| 149 |
rewards_list = cache
|
| 150 |
+
elif len(cache) > n:
|
| 151 |
+
rewards_list = cache[:n]
|
| 152 |
else:
|
| 153 |
+
# Still short β extend with mean rather than tile-mod so we don't
|
| 154 |
+
# duplicate rewards for the same prompt (tile-mod β zero advantage).
|
| 155 |
+
mean_r = sum(cache) / len(cache)
|
| 156 |
+
rewards_list = list(cache) + [mean_r] * (n - len(cache))
|
| 157 |
+
print(f"[PMOpsGRPOTrainer] WARNING: cache has {len(cache)} rewards for n={n}; "
|
| 158 |
+
f"padding with mean={mean_r:.3f}. Consider matching num_generations.")
|
| 159 |
|
| 160 |
device = self.accelerator.device
|
| 161 |
rewards = torch.tensor(
|
|
|
|
| 164 |
device=device,
|
| 165 |
).unsqueeze(1) # [batch_size, 1]
|
| 166 |
|
| 167 |
+
std = rewards.std().item() if n > 1 else 0.0
|
| 168 |
+
self.log({"reward/injected_mean": rewards.mean().item(),
|
| 169 |
+
"reward/injected_std": std})
|
| 170 |
self._rollout_reward_cache = [] # consume cache
|
| 171 |
+
print(f"[PMOpsGRPOTrainer] injected {n} rewards "
|
| 172 |
+
f"mean={rewards.mean().item():.3f} std={std:.3f}")
|
| 173 |
return rewards
|
| 174 |
|
| 175 |
if self.rollout_func is not None and self._rollout_capture_calls == 0 and not self._warned_rollout_bypass:
|
|
@@ -42,7 +42,8 @@ def _get_model_for_generation(trainer):
|
|
| 42 |
return trainer.model
|
| 43 |
|
| 44 |
|
| 45 |
-
def _generate_no_vllm(trainer, prompt_text: str, tokenizer,
|
|
|
|
| 46 |
"""Generate one completion using HF model.generate() without vLLM.
|
| 47 |
|
| 48 |
Returns the same dict shape as generate_rollout_completions so the rest
|
|
@@ -68,7 +69,9 @@ def _generate_no_vllm(trainer, prompt_text: str, tokenizer, max_new_tokens: int
|
|
| 68 |
**enc,
|
| 69 |
max_new_tokens=max_new_tokens,
|
| 70 |
do_sample=True,
|
| 71 |
-
temperature=
|
|
|
|
|
|
|
| 72 |
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
|
| 73 |
output_scores=True,
|
| 74 |
return_dict_in_generate=True,
|
|
@@ -233,15 +236,18 @@ def rollout_once(
|
|
| 233 |
tokenizer,
|
| 234 |
dataset_prompt: str,
|
| 235 |
max_steps: int = 15,
|
|
|
|
| 236 |
) -> dict:
|
| 237 |
"""Play one full PM-Ops episode. Returns trajectory + reward signals.
|
| 238 |
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
The env's hard limit (MAX_STEPS=40) still applies server-side.
|
| 242 |
"""
|
| 243 |
seed = parse_seed_from_prompt(dataset_prompt)
|
| 244 |
-
|
|
|
|
|
|
|
|
|
|
| 245 |
|
| 246 |
obs = result.observation if hasattr(result, "observation") else result
|
| 247 |
obs_dict = _obs_to_dict(obs)
|
|
@@ -262,6 +268,7 @@ def rollout_once(
|
|
| 262 |
final_score = 0.0
|
| 263 |
step = 0
|
| 264 |
done = False
|
|
|
|
| 265 |
|
| 266 |
# Channel tracking for reward_no_wrong_channels
|
| 267 |
oncall_channels: set[str] = set() # populated after reading runbook
|
|
@@ -324,11 +331,15 @@ def rollout_once(
|
|
| 324 |
new_obs = result.observation if hasattr(result, "observation") else result
|
| 325 |
obs_dict = _obs_to_dict(new_obs)
|
| 326 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 327 |
# Extract oncall channels from runbook response (available one step later)
|
| 328 |
if action_type == "meta.read_runbook":
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
data = last.get("data") or {}
|
| 332 |
if isinstance(data, dict):
|
| 333 |
org = data.get("org_config") or {}
|
| 334 |
oncall_channels = set(org.get("oncall_channels", {}).values())
|
|
@@ -356,30 +367,37 @@ def rollout_once(
|
|
| 356 |
|
| 357 |
read_runbook_reward = 1.0 if read_runbook_done else 0.0
|
| 358 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 359 |
# --- Reward gating ---
|
| 360 |
# If model never output valid JSON, it never actually tried anything.
|
| 361 |
# Strip all process rewards and apply a harsh penalty.
|
| 362 |
-
# Process rewards (runbook, no_wrong, efficiency) only matter if model acted.
|
| 363 |
if valid_action_count == 0:
|
| 364 |
combined = -1.0
|
| 365 |
elif final_score == 0.0:
|
| 366 |
-
# Model tried (valid JSON) but task failed
|
|
|
|
|
|
|
|
|
|
| 367 |
combined = (
|
| 368 |
-
|
| 369 |
+ read_runbook_reward * 0.10
|
| 370 |
- 0.30 # hard penalty for zero task completion
|
| 371 |
)
|
| 372 |
else:
|
| 373 |
combined = (
|
| 374 |
-
final_score * 0.
|
|
|
|
| 375 |
+ no_wrong_channels * 0.15
|
| 376 |
-
+ valid_json_ratio * 0.
|
| 377 |
+ read_runbook_reward * 0.15
|
| 378 |
+ efficiency * 0.10
|
| 379 |
)
|
| 380 |
|
| 381 |
print(
|
| 382 |
-
f"[rollout] steps={step} final={final_score:.3f} "
|
| 383 |
f"json={valid_json_ratio:.2f} runbook={read_runbook_reward:.0f} "
|
| 384 |
f"no_wrong={no_wrong_channels:.2f} eff={efficiency:.2f} "
|
| 385 |
f"valid_acts={valid_action_count} β combined={combined:.3f}"
|
|
@@ -420,13 +438,20 @@ def make_rollout_func(sync_env, tokenizer, max_steps: int = 15):
|
|
| 420 |
"read_runbook_reward": [],
|
| 421 |
"efficiency_reward": [],
|
| 422 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 423 |
for prompt_text in prompts:
|
|
|
|
|
|
|
| 424 |
episode = rollout_once(
|
| 425 |
trainer=trainer,
|
| 426 |
sync_env=sync_env,
|
| 427 |
tokenizer=tokenizer,
|
| 428 |
dataset_prompt=prompt_text,
|
| 429 |
max_steps=max_steps,
|
|
|
|
| 430 |
)
|
| 431 |
for k in out:
|
| 432 |
out[k].append(episode[k])
|
|
|
|
| 42 |
return trainer.model
|
| 43 |
|
| 44 |
|
| 45 |
+
def _generate_no_vllm(trainer, prompt_text: str, tokenizer,
|
| 46 |
+
max_new_tokens: int = 512, temperature: float = 1.1) -> dict:
|
| 47 |
"""Generate one completion using HF model.generate() without vLLM.
|
| 48 |
|
| 49 |
Returns the same dict shape as generate_rollout_completions so the rest
|
|
|
|
| 69 |
**enc,
|
| 70 |
max_new_tokens=max_new_tokens,
|
| 71 |
do_sample=True,
|
| 72 |
+
temperature=temperature,
|
| 73 |
+
top_p=0.95,
|
| 74 |
+
top_k=50,
|
| 75 |
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
|
| 76 |
output_scores=True,
|
| 77 |
return_dict_in_generate=True,
|
|
|
|
| 236 |
tokenizer,
|
| 237 |
dataset_prompt: str,
|
| 238 |
max_steps: int = 15,
|
| 239 |
+
gen_offset: int = 0,
|
| 240 |
) -> dict:
|
| 241 |
"""Play one full PM-Ops episode. Returns trajectory + reward signals.
|
| 242 |
|
| 243 |
+
gen_offset: added to seed so each GRPO generation explores a different env
|
| 244 |
+
episode even when receiving the same prompt (same base seed).
|
|
|
|
| 245 |
"""
|
| 246 |
seed = parse_seed_from_prompt(dataset_prompt)
|
| 247 |
+
if seed is not None:
|
| 248 |
+
result = sync_env.reset(seed=seed + gen_offset)
|
| 249 |
+
else:
|
| 250 |
+
result = sync_env.reset()
|
| 251 |
|
| 252 |
obs = result.observation if hasattr(result, "observation") else result
|
| 253 |
obs_dict = _obs_to_dict(obs)
|
|
|
|
| 268 |
final_score = 0.0
|
| 269 |
step = 0
|
| 270 |
done = False
|
| 271 |
+
ok_count = 0 # number of steps where env accepted the action (last_action_result.ok)
|
| 272 |
|
| 273 |
# Channel tracking for reward_no_wrong_channels
|
| 274 |
oncall_channels: set[str] = set() # populated after reading runbook
|
|
|
|
| 331 |
new_obs = result.observation if hasattr(result, "observation") else result
|
| 332 |
obs_dict = _obs_to_dict(new_obs)
|
| 333 |
|
| 334 |
+
# Track per-step env acceptance for intermediate reward signal
|
| 335 |
+
last_result = obs_dict.get("last_action_result") or {}
|
| 336 |
+
if last_result.get("ok"):
|
| 337 |
+
ok_count += 1
|
| 338 |
+
|
| 339 |
# Extract oncall channels from runbook response (available one step later)
|
| 340 |
if action_type == "meta.read_runbook":
|
| 341 |
+
if last_result.get("ok"):
|
| 342 |
+
data = last_result.get("data") or {}
|
|
|
|
| 343 |
if isinstance(data, dict):
|
| 344 |
org = data.get("org_config") or {}
|
| 345 |
oncall_channels = set(org.get("oncall_channels", {}).values())
|
|
|
|
| 367 |
|
| 368 |
read_runbook_reward = 1.0 if read_runbook_done else 0.0
|
| 369 |
|
| 370 |
+
# ok_ratio: fraction of steps where env accepted the action β provides per-step
|
| 371 |
+
# intermediate signal that varies across rollouts even when final_score=0.
|
| 372 |
+
ok_ratio = ok_count / max(step, 1)
|
| 373 |
+
|
| 374 |
# --- Reward gating ---
|
| 375 |
# If model never output valid JSON, it never actually tried anything.
|
| 376 |
# Strip all process rewards and apply a harsh penalty.
|
|
|
|
| 377 |
if valid_action_count == 0:
|
| 378 |
combined = -1.0
|
| 379 |
elif final_score == 0.0:
|
| 380 |
+
# Model tried (valid JSON) but task failed.
|
| 381 |
+
# ok_ratio gives gradient signal that distinguishes "wrong but env-valid" actions
|
| 382 |
+
# from "env-rejected" ones β this is the variance source GRPO needs when
|
| 383 |
+
# final_score is always 0 (model not yet solving tasks).
|
| 384 |
combined = (
|
| 385 |
+
ok_ratio * 0.20
|
| 386 |
+ read_runbook_reward * 0.10
|
| 387 |
- 0.30 # hard penalty for zero task completion
|
| 388 |
)
|
| 389 |
else:
|
| 390 |
combined = (
|
| 391 |
+
final_score * 0.40
|
| 392 |
+
+ ok_ratio * 0.10
|
| 393 |
+ no_wrong_channels * 0.15
|
| 394 |
+
+ valid_json_ratio * 0.10
|
| 395 |
+ read_runbook_reward * 0.15
|
| 396 |
+ efficiency * 0.10
|
| 397 |
)
|
| 398 |
|
| 399 |
print(
|
| 400 |
+
f"[rollout] steps={step} final={final_score:.3f} ok={ok_ratio:.2f} "
|
| 401 |
f"json={valid_json_ratio:.2f} runbook={read_runbook_reward:.0f} "
|
| 402 |
f"no_wrong={no_wrong_channels:.2f} eff={efficiency:.2f} "
|
| 403 |
f"valid_acts={valid_action_count} β combined={combined:.3f}"
|
|
|
|
| 438 |
"read_runbook_reward": [],
|
| 439 |
"efficiency_reward": [],
|
| 440 |
}
|
| 441 |
+
# Track how many times each unique prompt has appeared so we can pass
|
| 442 |
+
# a gen_offset β ensures repeated prompts (num_generations > 1) hit
|
| 443 |
+
# different env seeds and produce different rollouts.
|
| 444 |
+
prompt_seen: dict[str, int] = {}
|
| 445 |
for prompt_text in prompts:
|
| 446 |
+
gen_offset = prompt_seen.get(prompt_text, 0)
|
| 447 |
+
prompt_seen[prompt_text] = gen_offset + 1
|
| 448 |
episode = rollout_once(
|
| 449 |
trainer=trainer,
|
| 450 |
sync_env=sync_env,
|
| 451 |
tokenizer=tokenizer,
|
| 452 |
dataset_prompt=prompt_text,
|
| 453 |
max_steps=max_steps,
|
| 454 |
+
gen_offset=gen_offset,
|
| 455 |
)
|
| 456 |
for k in out:
|
| 457 |
out[k].append(episode[k])
|
|
@@ -512,106 +512,7 @@
|
|
| 512 |
"id": "cell-14",
|
| 513 |
"metadata": {},
|
| 514 |
"outputs": [],
|
| 515 |
-
"source": [
|
| 516 |
-
"from trl.experimental.openenv import generate_rollout_completions\n",
|
| 517 |
-
"from training.rollout import (\n",
|
| 518 |
-
" _obs_to_dict, _current_obs_text, build_messages,\n",
|
| 519 |
-
" extract_json_action, step_aware_fallback,\n",
|
| 520 |
-
")\n",
|
| 521 |
-
"from training.dataset import parse_seed_from_prompt\n",
|
| 522 |
-
"\n",
|
| 523 |
-
"grpo_env = GenericEnvClient(base_url=ENV_URL).sync()\n",
|
| 524 |
-
"grpo_env.connect()\n",
|
| 525 |
-
"print('GRPO training env connected')\n",
|
| 526 |
-
"\n",
|
| 527 |
-
"\n",
|
| 528 |
-
"def run_grpo_episode(trainer, env, tok, dataset_prompt, max_steps=TRAIN_MAX_STEPS):\n",
|
| 529 |
-
" \"\"\"Run one full PM-ops episode and return flat trajectory + reward.\"\"\"\n",
|
| 530 |
-
" seed = parse_seed_from_prompt(dataset_prompt)\n",
|
| 531 |
-
" result = env.reset(seed=seed) if seed is not None else env.reset()\n",
|
| 532 |
-
" obs_dict = _obs_to_dict(result.observation if hasattr(result, 'observation') else result)\n",
|
| 533 |
-
" task_brief = obs_dict.get('task_brief') or dataset_prompt\n",
|
| 534 |
-
"\n",
|
| 535 |
-
" prompt_ids, completion_ids, logprobs = [], [], []\n",
|
| 536 |
-
" turn_history = []\n",
|
| 537 |
-
" valid_json_count = 0\n",
|
| 538 |
-
" env_score = 0.0\n",
|
| 539 |
-
" step, done = 0, False\n",
|
| 540 |
-
" _sample_logged = False\n",
|
| 541 |
-
"\n",
|
| 542 |
-
" while not done and step < max_steps:\n",
|
| 543 |
-
" obs_text = _current_obs_text(obs_dict, step, task_brief)\n",
|
| 544 |
-
" msgs = build_messages(turn_history, obs_text)\n",
|
| 545 |
-
" prompt_text = tok.apply_chat_template(\n",
|
| 546 |
-
" msgs, add_generation_prompt=True, tokenize=False, enable_thinking=False\n",
|
| 547 |
-
" )\n",
|
| 548 |
-
"\n",
|
| 549 |
-
" rollout_out = generate_rollout_completions(trainer, [prompt_text])[0]\n",
|
| 550 |
-
" prompt_ids.extend(rollout_out['prompt_ids'])\n",
|
| 551 |
-
" completion_ids.extend(rollout_out['completion_ids'])\n",
|
| 552 |
-
" logprobs.extend(rollout_out['logprobs'])\n",
|
| 553 |
-
"\n",
|
| 554 |
-
" completion_text = rollout_out.get('text') or tok.decode(\n",
|
| 555 |
-
" rollout_out['completion_ids'], skip_special_tokens=True\n",
|
| 556 |
-
" )\n",
|
| 557 |
-
"\n",
|
| 558 |
-
" # Log one sample per episode so we can visually track format quality\n",
|
| 559 |
-
" if not _sample_logged:\n",
|
| 560 |
-
" print(f' [sample] {repr(completion_text[:180])}')\n",
|
| 561 |
-
" _sample_logged = True\n",
|
| 562 |
-
"\n",
|
| 563 |
-
" parsed = extract_json_action(completion_text)\n",
|
| 564 |
-
" if parsed is not None:\n",
|
| 565 |
-
" valid_json_count += 1\n",
|
| 566 |
-
" else:\n",
|
| 567 |
-
" parsed = step_aware_fallback(step, max_steps)\n",
|
| 568 |
-
"\n",
|
| 569 |
-
" action_type = parsed.get('action_type', 'meta.noop')\n",
|
| 570 |
-
" args = parsed.get('args', {})\n",
|
| 571 |
-
"\n",
|
| 572 |
-
" turn_history.append({\n",
|
| 573 |
-
" 'obs_text' : obs_text,\n",
|
| 574 |
-
" 'completion': completion_text,\n",
|
| 575 |
-
" 'is_runbook': (action_type == 'meta.read_runbook' and parsed is not None),\n",
|
| 576 |
-
" })\n",
|
| 577 |
-
"\n",
|
| 578 |
-
" result = env.step({'action_type': action_type, 'args': args})\n",
|
| 579 |
-
" obs_dict = _obs_to_dict(result.observation if hasattr(result, 'observation') else result)\n",
|
| 580 |
-
" done = bool(getattr(result, 'done', obs_dict.get('done', False)))\n",
|
| 581 |
-
" env_score = float(getattr(result, 'reward', obs_dict.get('reward', 0.0)))\n",
|
| 582 |
-
" step += 1\n",
|
| 583 |
-
"\n",
|
| 584 |
-
" json_ratio = valid_json_count / max(step, 1)\n",
|
| 585 |
-
" reward = env_score * 0.85 + json_ratio * 0.15\n",
|
| 586 |
-
" print(f' [rollout] steps={step} env={env_score:.3f} json={json_ratio:.2f} -> reward={reward:.3f}')\n",
|
| 587 |
-
" return {\n",
|
| 588 |
-
" 'prompt_ids' : prompt_ids,\n",
|
| 589 |
-
" 'completion_ids': completion_ids,\n",
|
| 590 |
-
" 'logprobs' : logprobs,\n",
|
| 591 |
-
" 'reward' : reward,\n",
|
| 592 |
-
" }\n",
|
| 593 |
-
"\n",
|
| 594 |
-
"\n",
|
| 595 |
-
"def grpo_rollout_func(prompts, trainer=None):\n",
|
| 596 |
-
" out = {'prompt_ids': [], 'completion_ids': [], 'logprobs': [], 'reward': []}\n",
|
| 597 |
-
" for prompt in prompts:\n",
|
| 598 |
-
" ep = run_grpo_episode(trainer, grpo_env, tokenizer, prompt, TRAIN_MAX_STEPS)\n",
|
| 599 |
-
" for k in out:\n",
|
| 600 |
-
" out[k].append(ep[k])\n",
|
| 601 |
-
" return out\n",
|
| 602 |
-
"\n",
|
| 603 |
-
"\n",
|
| 604 |
-
"def grpo_reward_func(completions, **kwargs):\n",
|
| 605 |
-
" \"\"\"Passthrough β reward is pre-computed in grpo_rollout_func.\"\"\"\n",
|
| 606 |
-
" rewards = kwargs.get('reward', [])\n",
|
| 607 |
-
" if not rewards:\n",
|
| 608 |
-
" print(f'[ERROR] reward_func: no reward key. kwargs keys: {list(kwargs.keys())}')\n",
|
| 609 |
-
" return [0.0] * len(completions)\n",
|
| 610 |
-
" return [float(r) for r in rewards]\n",
|
| 611 |
-
"\n",
|
| 612 |
-
"\n",
|
| 613 |
-
"print(f'GRPO rollout ready max_steps={TRAIN_MAX_STEPS}')"
|
| 614 |
-
]
|
| 615 |
},
|
| 616 |
{
|
| 617 |
"cell_type": "markdown",
|
|
@@ -917,4 +818,4 @@
|
|
| 917 |
},
|
| 918 |
"nbformat": 4,
|
| 919 |
"nbformat_minor": 5
|
| 920 |
-
}
|
|
|
|
| 512 |
"id": "cell-14",
|
| 513 |
"metadata": {},
|
| 514 |
"outputs": [],
|
| 515 |
+
"source": "from training.rollout import (\n _obs_to_dict, _current_obs_text, build_messages,\n extract_json_action, step_aware_fallback,\n _generate_no_vllm,\n)\nfrom training.dataset import parse_seed_from_prompt\n\ngrpo_env = GenericEnvClient(base_url=ENV_URL).sync()\ngrpo_env.connect()\nprint('GRPO training env connected')\n\n# Must be > 1.0 so rollouts of the same prompt diverge and GRPO gets non-zero advantage.\nROLLOUT_TEMPERATURE = 1.1\n\n\ndef run_grpo_episode(trainer, env, tok, dataset_prompt, max_steps=TRAIN_MAX_STEPS,\n gen_offset=0):\n \"\"\"Run one PM-ops episode. gen_offset shifts the env seed so each GRPO\n generation for the same prompt explores a different env episode.\n \"\"\"\n seed = parse_seed_from_prompt(dataset_prompt)\n if seed is not None:\n result = env.reset(seed=seed + gen_offset)\n else:\n result = env.reset()\n\n obs_dict = _obs_to_dict(result.observation if hasattr(result, 'observation') else result)\n task_brief = obs_dict.get('task_brief') or dataset_prompt\n\n prompt_ids, completion_ids, logprobs = [], [], []\n turn_history = []\n valid_json_count = 0\n ok_count = 0 # steps where env accepted the action (last_action_result.ok)\n read_runbook_done = False\n env_score = 0.0\n step, done = 0, False\n _sample_logged = False\n\n while not done and step < max_steps:\n obs_text = _current_obs_text(obs_dict, step, task_brief)\n msgs = build_messages(turn_history, obs_text)\n prompt_text = tok.apply_chat_template(\n msgs, add_generation_prompt=True, tokenize=False, enable_thinking=False\n )\n\n # Always use _generate_no_vllm at ROLLOUT_TEMPERATURE.\n # generate_rollout_completions defaults to greedy/low-temp, which collapses\n # all rollouts for the same prompt to identical outputs β zero GRPO advantage.\n rollout_out = _generate_no_vllm(\n trainer, prompt_text, tok,\n max_new_tokens=MAX_COMP_LEN,\n temperature=ROLLOUT_TEMPERATURE,\n )\n prompt_ids.extend(rollout_out['prompt_ids'])\n completion_ids.extend(rollout_out['completion_ids'])\n logprobs.extend(rollout_out['logprobs'])\n completion_text = rollout_out['text']\n\n if not _sample_logged:\n print(f' [sample] {repr(completion_text[:180])}')\n _sample_logged = True\n\n parsed = extract_json_action(completion_text)\n if parsed is not None:\n valid_json_count += 1\n else:\n parsed = step_aware_fallback(step, max_steps)\n\n action_type = parsed.get('action_type', 'meta.noop')\n args = parsed.get('args', {})\n\n if action_type == 'meta.read_runbook':\n read_runbook_done = True\n\n turn_history.append({\n 'obs_text' : obs_text,\n 'completion': completion_text,\n 'is_runbook': (action_type == 'meta.read_runbook' and parsed is not None),\n })\n\n try:\n result = env.step({'action_type': action_type, 'args': args})\n except RuntimeError as exc:\n if 'VALIDATION_ERROR' in str(exc):\n result = env.step({'action_type': 'meta.noop', 'args': {}})\n else:\n raise\n\n obs_dict = _obs_to_dict(result.observation if hasattr(result, 'observation') else result)\n\n # Per-step intermediate signal: did env accept this action?\n last_res = obs_dict.get('last_action_result') or {}\n if last_res.get('ok'):\n ok_count += 1\n\n done = bool(getattr(result, 'done', obs_dict.get('done', False)))\n env_score = float(getattr(result, 'reward', obs_dict.get('reward', 0.0)))\n step += 1\n\n json_ratio = valid_json_count / max(step, 1)\n ok_ratio = ok_count / max(step, 1)\n rb_reward = 1.0 if read_runbook_done else 0.0\n\n # Reward formula designed to produce VARIANCE across rollouts:\n # - ok_ratio varies per episode (different actions accepted/rejected by env)\n # - env_score varies when the model solves the task with different quality\n # Using json_ratio alone (always ~1.0 after SFT) gives a constant reward β zero gradient.\n if valid_json_count == 0:\n reward = -1.0\n elif env_score == 0.0:\n # Task not completed β ok_ratio differentiates rollouts that took env-valid\n # actions from ones that didn't, giving GRPO a gradient even before task completion.\n reward = ok_ratio * 0.20 + rb_reward * 0.10 - 0.30\n else:\n reward = env_score * 0.60 + ok_ratio * 0.10 + json_ratio * 0.10 + rb_reward * 0.20\n\n reward = max(-1.0, min(1.0, reward))\n print(f' [rollout] steps={step} env={env_score:.3f} ok={ok_ratio:.2f} '\n f'json={json_ratio:.2f} rb={rb_reward:.0f} offset={gen_offset} -> reward={reward:.3f}')\n return {\n 'prompt_ids' : prompt_ids,\n 'completion_ids': completion_ids,\n 'logprobs' : logprobs,\n 'reward' : reward,\n }\n\n\ndef grpo_rollout_func(prompts, trainer=None):\n out = {'prompt_ids': [], 'completion_ids': [], 'logprobs': [], 'reward': []}\n # Count how many times each prompt has appeared so far in this batch.\n # TRL sends [p, p, ...] (same prompt num_gen times) β gen_offset breaks symmetry.\n prompt_seen: dict = {}\n for prompt in prompts:\n gen_offset = prompt_seen.get(prompt, 0)\n prompt_seen[prompt] = gen_offset + 1\n ep = run_grpo_episode(trainer, grpo_env, tokenizer, prompt,\n TRAIN_MAX_STEPS, gen_offset=gen_offset)\n for k in out:\n out[k].append(ep[k])\n return out\n\n\n_warned_missing_reward = False\ndef grpo_reward_func(completions, **kwargs):\n \"\"\"Passthrough β reward is pre-computed in grpo_rollout_func.\"\"\"\n global _warned_missing_reward\n rewards = kwargs.get('reward', [])\n if not rewards:\n if not _warned_missing_reward:\n print(f\"[WARN] reward_func fallback with no reward key. kwargs keys: {list(kwargs.keys())}\")\n _warned_missing_reward = True\n return [0.0] * len(completions)\n return [float(r) for r in rewards]\n\n\nprint(f'GRPO rollout ready max_steps={TRAIN_MAX_STEPS} temperature={ROLLOUT_TEMPERATURE}')"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 516 |
},
|
| 517 |
{
|
| 518 |
"cell_type": "markdown",
|
|
|
|
| 818 |
},
|
| 819 |
"nbformat": 4,
|
| 820 |
"nbformat_minor": 5
|
| 821 |
+
}
|