import ast import json import multiprocessing as mp import re from dataclasses import dataclass, field from typing import Any, Optional import trackio from datasets import load_dataset from transformers import AutoTokenizer, HfArgumentParser, TrainerCallback from trl import GRPOConfig, GRPOTrainer SYSTEM_PROMPT = ( "You are a Python code generator. Return only Python code. " "Write a correct solution function for the task. Include a concise but proper docstring on the main function." ) @dataclass class ScriptArgs: model_name_or_path: str = field(default="Qwen/Qwen3-1.7B-Base") dataset_name: str = field(default="google-research-datasets/mbpp") dataset_config: str = field(default="sanitized") train_split: str = field(default="train") eval_split: str = field(default="validation") max_train_samples: Optional[int] = field(default=300) max_eval_samples: Optional[int] = field(default=64) attn_implementation: str = field(default="sdpa") reward_timeout: int = field(default=4) push_to_hub: bool = field(default=True) hub_model_id: str = field(default="AbhilekhMeda/qwen3-1.7b-grpo-python-mbpp") run_name: str = field(default="grpo_qwen3_1p7b_mbpp_exec_reward") project: str = field(default="grpo-qwen3-python-code") trackio_space_id: str = field(default="AbhilekhMeda/mlintern-grpoqwen") def _strip_code_fence(text: str) -> str: text = text.strip() match = re.search(r"```(?:python)?\n(.*?)```", text, re.DOTALL | re.IGNORECASE) return match.group(1).strip() if match else text def _extract_function_name_from_tests(test_list: list[str]) -> Optional[str]: for test in test_list: m = re.search(r"assert\s+([A-Za-z_][A-Za-z0-9_]*)\s*\(", test) if m: return m.group(1) return None def _build_prompt(example: dict[str, Any]) -> dict[str, Any]: prompt_text = ( f"Task:\n{example['prompt']}\n\n" "Requirements:\n" "1. Return only Python code.\n" "2. Implement the requested function.\n" "3. Include a proper docstring on the main function.\n" "4. Do not print example usage.\n" ) return { "prompt": [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": prompt_text}, ], "task_id": example["task_id"], "raw_prompt": example["prompt"], "test_imports": example["test_imports"], "test_list": example["test_list"], "reference_code": example["code"], } def _docstring_ok(code: str, fn_name: Optional[str]) -> bool: if not fn_name: return False try: tree = ast.parse(code) for node in tree.body: if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == fn_name: doc = ast.get_docstring(node) return bool(doc and len(doc.strip()) >= 10) except Exception: return False return False def _runner(queue, code: str, imports: list[str], tests: list[str]): glb: dict[str, Any] = {} try: exec(compile(code, "", "exec"), glb, glb) if imports: exec("\n".join(imports), glb, glb) for test in tests: exec(test, glb, glb) queue.put({"ran": True, "passed": True, "error": ""}) except Exception as e: queue.put({"ran": True, "passed": False, "error": repr(e)}) def _exec_in_subprocess(code: str, test_imports: list[str], tests: list[str], timeout: int) -> dict[str, Any]: ctx = mp.get_context("spawn") queue = ctx.Queue() proc = ctx.Process(target=_runner, args=(queue, code, test_imports, tests)) proc.start() proc.join(timeout) if proc.is_alive(): proc.terminate() proc.join(1) return {"ran": False, "passed": False, "error": "timeout"} return queue.get() if not queue.empty() else {"ran": False, "passed": False, "error": "no_result"} def execution_reward(completions, test_imports, test_list, log_extra=None, log_metric=None, **kwargs): rewards = [] run_rates, pass_rates, doc_rates, errors = [], [], [], [] for completion, imports, tests in zip(completions, test_imports, test_list): content = completion[0]["content"] if isinstance(completion, list) else completion code = _strip_code_fence(content) fn_name = _extract_function_name_from_tests(tests) result = _exec_in_subprocess(code, imports, tests, timeout=4) ran = 1.0 if result["ran"] and result["error"] != "timeout" else 0.0 passed = 1.0 if result["passed"] else 0.0 doc = 1.0 if _docstring_ok(code, fn_name) else 0.0 rewards.append(0.25 * ran + 0.6 * passed + 0.15 * doc) run_rates.append(ran) pass_rates.append(passed) doc_rates.append(doc) errors.append(result["error"][:120]) if log_extra: log_extra("exec_error", errors) log_extra("passed_tests", [str(x) for x in pass_rates]) log_extra("docstring_ok", [str(x) for x in doc_rates]) if log_metric and rewards: n = float(len(rewards)) log_metric("exec_run_rate", sum(run_rates) / n) log_metric("exec_pass_rate", sum(pass_rates) / n) log_metric("docstring_rate", sum(doc_rates) / n) return [float(r) for r in rewards] class AlertCallback(TrainerCallback): def on_log(self, args, state, control, logs=None, **kwargs): if not logs: return if logs.get("loss") is not None and logs["loss"] > 2.0: trackio.alert("high_loss", f"loss={logs['loss']:.4f} at step {state.global_step} — try lr x0.5 if it persists", level="WARN") if logs.get("reward") is not None and logs["reward"] > 0.8: trackio.alert("strong_reward", f"reward={logs['reward']:.4f} at step {state.global_step} — keep current config and refine around lr", level="INFO") if logs.get("completions/clipped_ratio") is not None and logs["completions/clipped_ratio"] > 0.3: trackio.alert("high_clipping", f"clipped_ratio={logs['completions/clipped_ratio']:.4f} at step {state.global_step} — increase max_completion_length if many outputs are truncated", level="WARN") def main(): parser = HfArgumentParser((ScriptArgs, GRPOConfig)) script_args, training_args = parser.parse_args_into_dataclasses() dataset = load_dataset(script_args.dataset_name, script_args.dataset_config) train_dataset = dataset[script_args.train_split] eval_dataset = dataset[script_args.eval_split] if script_args.max_train_samples: train_dataset = train_dataset.select(range(min(script_args.max_train_samples, len(train_dataset)))) if script_args.max_eval_samples: eval_dataset = eval_dataset.select(range(min(script_args.max_eval_samples, len(eval_dataset)))) train_dataset = train_dataset.map(_build_prompt, remove_columns=train_dataset.column_names) eval_dataset = eval_dataset.map(_build_prompt, remove_columns=eval_dataset.column_names) tokenizer = AutoTokenizer.from_pretrained(script_args.model_name_or_path, padding_side="left") if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token training_args.report_to = "trackio" training_args.run_name = script_args.run_name training_args.project = script_args.project training_args.trackio_space_id = script_args.trackio_space_id training_args.push_to_hub = script_args.push_to_hub training_args.hub_model_id = script_args.hub_model_id training_args.disable_tqdm = True training_args.logging_strategy = "steps" training_args.logging_first_step = True training_args.save_strategy = "steps" training_args.eval_strategy = "steps" training_args.remove_unused_columns = False training_args.model_init_kwargs = { "attn_implementation": script_args.attn_implementation, "torch_dtype": "bfloat16", } training_args.chat_template_kwargs = {"enable_thinking": False} trainer = GRPOTrainer( model=script_args.model_name_or_path, args=training_args, processing_class=tokenizer, reward_funcs=[execution_reward], train_dataset=train_dataset, eval_dataset=eval_dataset, callbacks=[AlertCallback()], ) trackio.alert("run_start", f"Starting GRPO code run with train_samples={len(train_dataset)} eval_samples={len(eval_dataset)} lr={training_args.learning_rate}", level="INFO") trainer.train(resume_from_checkpoint=training_args.resume_from_checkpoint) metrics = trainer.evaluate() trainer.log_metrics("eval", metrics) trainer.save_metrics("eval", metrics) trainer.save_model(training_args.output_dir) if training_args.push_to_hub: trainer.push_to_hub(commit_message="End of GRPO training") trackio.alert("run_complete", f"Training complete at step {trainer.state.global_step} with eval metrics: {json.dumps(metrics, default=str)}", level="INFO") if __name__ == "__main__": main()