File size: 2,379 Bytes
acf77ab
 
 
a38ca5e
 
 
 
 
 
 
 
 
 
dbf2d43
 
 
acf77ab
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
import sys
import os

# ── Compatibility shim ──────────────────────────────────────────────────────
# torchao>=0.9 decorates enums with @register_as_pytree_constant which calls
# torch.utils._pytree.register_constant – a function that only exists in
# PyTorch 2.7+.  We're running on 2.6.0, so patch it in before any import
# of torchao / transformers / unsloth triggers the missing attribute error.
import torch.utils._pytree as _pytree
if not hasattr(_pytree, "register_constant"):
    _pytree.register_constant = lambda cls: cls  # no-op shim
# ────────────────────────────────────────────────────────────────────────────

os.environ["WANDB_API_KEY"] = "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"
os.environ["WANDB_PROJECT"] = "codeforge-grpo"

# Add the root directory to PYTHONPATH so codeforge module can be imported
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()