SavK1 Claude Sonnet 4.6 commited on
Commit
38457df
Β·
1 Parent(s): fff1405

fix(grpo): eliminate zero-advantage collapse from stale rewards

Browse files

Three 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 CHANGED
@@ -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 once here
134
- # to populate reward cache from the exact prompt batch.
 
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
- rewards_list = [cache[i % len(cache)] for i in range(n)]
 
 
 
 
 
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
- self.log({"reward/injected_mean": rewards.mean().item()})
 
 
161
  self._rollout_reward_cache = [] # consume cache
162
- print(f"[PMOpsGRPOTrainer] injected {n} rewards, mean={rewards.mean().item():.3f}")
 
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:
training/rollout.py CHANGED
@@ -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, max_new_tokens: int = 512) -> dict:
 
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=0.7,
 
 
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
- max_steps: training cap per task type. Triage can be solved in 5 steps;
240
- capping at 15 avoids burning compute on aimless late-episode steps.
241
- The env's hard limit (MAX_STEPS=40) still applies server-side.
242
  """
243
  seed = parse_seed_from_prompt(dataset_prompt)
244
- result = sync_env.reset(seed=seed) if seed is not None else sync_env.reset()
 
 
 
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
- last = obs_dict.get("last_action_result") or {}
330
- if last.get("ok"):
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 β€” small process credit, no final bonus
 
 
 
367
  combined = (
368
- valid_json_ratio * 0.15
369
  + read_runbook_reward * 0.10
370
  - 0.30 # hard penalty for zero task completion
371
  )
372
  else:
373
  combined = (
374
- final_score * 0.45
 
375
  + no_wrong_channels * 0.15
376
- + valid_json_ratio * 0.15
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])
training/train_v3.ipynb CHANGED
@@ -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
+ }