Spaces:
Sleeping
Sleeping
Fix overnight training crash on empty completions
Browse files- 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=
|
| 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":
|
| 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),
|