| import sys |
| import os |
|
|
| |
| |
| |
| |
| |
| import torch.utils._pytree as _pytree |
| if not hasattr(_pytree, "register_constant"): |
| _pytree.register_constant = lambda cls: cls |
| |
|
|
| os.environ["WANDB_API_KEY"] = "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" |
| os.environ["WANDB_PROJECT"] = "codeforge-grpo" |
|
|
| |
| sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) |
|
|
| import argparse |
| from trainer.model_boot import load_model |
| from trainer.grpo_config import build_config |
| from trainer.reward_fn import codeforge_reward_fn |
| from trainer.mbpp_adapter import load_mbpp_dataset |
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--steps", type=int, default=300) |
| args = parser.parse_args() |
|
|
| print("Loading model...") |
| model, tokenizer = load_model() |
|
|
| config = build_config( |
| run_name="codeforge-1.5b", |
| output_dir="checkpoints/codeforge_1.5b", |
| ) |
| config.max_steps = args.steps |
|
|
| print("Loading dataset...") |
| dataset = load_mbpp_dataset(split="train") |
|
|
| print("Initializing trainer...") |
| from trl import GRPOTrainer |
| trainer = GRPOTrainer( |
| model=model, |
| processing_class=tokenizer, |
| reward_funcs=[codeforge_reward_fn], |
| args=config, |
| train_dataset=dataset, |
| ) |
| |
| print("Starting training!") |
| trainer.train() |
| |
| print("Training complete, saving...") |
| model.save_pretrained("checkpoints/codeforge_1.5b_final") |
| print("Saved to checkpoints/codeforge_1.5b_final") |
|
|
| if __name__ == "__main__": |
| main() |
|
|