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()
|