Dishaaa25 commited on
Commit
de4b527
·
1 Parent(s): a8426bb

Fix overnight training crash on empty completions

Browse files
Files changed (1) hide show
  1. training/train_grpo.py +11 -2
training/train_grpo.py CHANGED
@@ -161,6 +161,14 @@ def extract_code(completion: str) -> str:
161
  return text
162
 
163
 
 
 
 
 
 
 
 
 
164
  def format_examples(problem: dict[str, Any]) -> str:
165
  visible_cases = [test_case for test_case in problem.get("test_cases", []) if test_case.get("is_visible", False)]
166
  if not visible_cases:
@@ -640,6 +648,7 @@ def build_reward_func(
640
  rewards: list[float] = []
641
 
642
  for prompt, completion in zip(prompts, completions):
 
643
  problem = controller.resolve_prompt(prompt)
644
  env = AdaptEnvironment(generator=controller.generator, generator_mode=controller.mode)
645
  env.reset(
@@ -651,7 +660,7 @@ def build_reward_func(
651
  observation = env.step(
652
  AdaptAction(
653
  session_id=env.session_id,
654
- code=extract_code(completion),
655
  )
656
  )
657
  rewards.append(float(observation.reward))
@@ -677,7 +686,7 @@ def build_reward_func(
677
  "problem_id": problem["problem_id"],
678
  "teacher_prompt": prompt,
679
  "solver_completion": completion,
680
- "extracted_code": extract_code(completion),
681
  "feedback": observation.feedback,
682
  "efficiency_score": observation.reward_components.get("efficiency_score"),
683
  "optimization_hints": extract_optimization_hints(observation.feedback),
 
161
  return text
162
 
163
 
164
+ def extract_code_for_execution(completion: str) -> str:
165
+ extracted = extract_code(completion)
166
+ if extracted.strip():
167
+ return extracted
168
+ # Empty generations should count as bad samples, not crash the whole run.
169
+ return "pass"
170
+
171
+
172
  def format_examples(problem: dict[str, Any]) -> str:
173
  visible_cases = [test_case for test_case in problem.get("test_cases", []) if test_case.get("is_visible", False)]
174
  if not visible_cases:
 
648
  rewards: list[float] = []
649
 
650
  for prompt, completion in zip(prompts, completions):
651
+ extracted_code = extract_code_for_execution(completion)
652
  problem = controller.resolve_prompt(prompt)
653
  env = AdaptEnvironment(generator=controller.generator, generator_mode=controller.mode)
654
  env.reset(
 
660
  observation = env.step(
661
  AdaptAction(
662
  session_id=env.session_id,
663
+ code=extracted_code,
664
  )
665
  )
666
  rewards.append(float(observation.reward))
 
686
  "problem_id": problem["problem_id"],
687
  "teacher_prompt": prompt,
688
  "solver_completion": completion,
689
+ "extracted_code": extracted_code,
690
  "feedback": observation.feedback,
691
  "efficiency_score": observation.reward_components.get("efficiency_score"),
692
  "optimization_hints": extract_optimization_hints(observation.feedback),