fathom-code / scripts /run_training.py
23f2002275
feat: phase 1 complete — smoke green on HF Jobs, training scripts, plot generator, Colab notebook, submission preflight
8787bd3
Raw
History Blame Contribute Delete
7.75 kB
"""FATHOM training orchestrator — runs inside GPU HF Space.
Pipeline: Dataset Gen -> SFT warm-start -> GRPO training -> Push to Hub
This script is the CMD entrypoint for the training Dockerfile.
"""
from __future__ import annotations
import json
import logging
import os
import subprocess
import sys
import time
from pathlib import Path
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
handlers=[logging.StreamHandler(sys.stdout)],
)
log = logging.getLogger("fathom.orchestrator")
# Config
HF_TOKEN = os.environ.get("HF_TOKEN", "")
MODEL_REPO = os.environ.get("MODEL_REPO", "Pratham-math/fathom-0.5b-grpo")
ENV_SPACE_URL = os.environ.get("ENV_URL", "https://Pratham-math-fathom-env.hf.space")
USE_SMOKE_MODEL = os.environ.get("USE_SMOKE_MODEL", "true").lower() == "true"
OUTPUT_DIR = Path("/app/outputs")
def _run(cmd: str, check: bool = True) -> int:
"""Run shell command with live output."""
log.info(">>> %s", cmd)
result = subprocess.run(cmd, shell=True, cwd="/app")
if check and result.returncode != 0:
log.error("Command failed with exit code %d", result.returncode)
return result.returncode
def step_1_verify_env():
"""Verify GPU + env server health."""
log.info("=" * 60)
log.info("STEP 1: Environment verification")
log.info("=" * 60)
# GPU check
import torch
assert torch.cuda.is_available(), "No GPU found!"
gpu_name = torch.cuda.get_device_name(0)
vram_gb = torch.cuda.get_device_properties(0).total_mem / 1e9
log.info("GPU: %s (%.1f GB VRAM)", gpu_name, vram_gb)
# Env server health
import urllib.request
try:
health_url = f"{ENV_SPACE_URL.rstrip('/')}/healthz"
r = urllib.request.urlopen(health_url, timeout=15)
body = json.loads(r.read().decode())
assert body.get("status") == "ok", f"Env health failed: {body}"
log.info("Env server healthy: %s", health_url)
except Exception as e:
log.warning("Env server not reachable (%s) — GRPO will use local env", e)
# HF Token
if HF_TOKEN:
log.info("HF_TOKEN present — will push model to %s", MODEL_REPO)
else:
log.warning("HF_TOKEN not set — model will be saved locally only")
def step_2_generate_dataset():
"""Generate deterministic dataset."""
log.info("=" * 60)
log.info("STEP 2: Dataset generation")
log.info("=" * 60)
from data.generate import generate_all
result = generate_all(output_dir=str(OUTPUT_DIR / "data"))
log.info(
"Dataset: train=%d eval=%d sft=%d",
result["train_count"],
result["eval_count"],
result["sft_count"],
)
return result
def step_3_sft_warmstart():
"""SFT warm-start training."""
log.info("=" * 60)
log.info("STEP 3: SFT warm-start")
log.info("=" * 60)
from hydra import initialize, compose
from train.model_load import load_model_and_tokenizer
from train.sft import run_sft
model_override = "model=qwen_0_5b_smoke" if USE_SMOKE_MODEL else "model=qwen_1_5b"
log.info("Using model config: %s", model_override)
with initialize(config_path="configs", version_base="1.3"):
cfg = compose(
config_name="config",
overrides=[
model_override,
"train=sft",
f"output_dir={OUTPUT_DIR}",
f"data.sft_traces_path={OUTPUT_DIR}/data/sft_traces.jsonl",
],
)
log.info("Loading model: %s", cfg.model.name)
model, tokenizer = load_model_and_tokenizer(cfg)
log.info("Starting SFT training...")
start = time.time()
adapter_dir = run_sft(cfg, model, tokenizer)
elapsed = time.time() - start
log.info("SFT complete in %.1f min. Adapter: %s", elapsed / 60, adapter_dir)
# Free GPU memory
del model, tokenizer
import torch
torch.cuda.empty_cache()
return adapter_dir
def step_4_grpo_training():
"""GRPO RL training."""
log.info("=" * 60)
log.info("STEP 4: GRPO training")
log.info("=" * 60)
from hydra import initialize, compose
from train.model_load import load_model_and_tokenizer
from train.grpo import run_grpo
from rewards.compose import make_reward_fn
from omegaconf import OmegaConf
model_override = "model=qwen_0_5b_smoke" if USE_SMOKE_MODEL else "model=qwen_1_5b"
with initialize(config_path="configs", version_base="1.3"):
cfg = compose(
config_name="config",
overrides=[
model_override,
"train=grpo",
f"output_dir={OUTPUT_DIR}",
"+hub.push=true",
f"+hub.repo_id={MODEL_REPO}",
],
)
log.info("Loading model for GRPO: %s", cfg.model.name)
model, tokenizer = load_model_and_tokenizer(cfg)
# Build reward function from config
cfg_reward = OmegaConf.create({
"alpha": float(cfg.reward.alpha),
"weights": OmegaConf.to_container(cfg.reward.weights, resolve=True),
"token_budget_variant": str(cfg.reward.token_budget_variant),
"max_calls": int(cfg.reward.max_calls),
})
reward_fn = make_reward_fn(cfg_reward)
log.info("Starting GRPO training (%d steps)...", cfg.train.max_steps)
start = time.time()
try:
merged_dir = run_grpo(
cfg, model, tokenizer, reward_fn,
env_url=ENV_SPACE_URL,
)
elapsed = time.time() - start
log.info("GRPO complete in %.1f min. Model: %s", elapsed / 60, merged_dir)
except Exception as e:
log.error("GRPO training failed: %s", e)
import traceback
traceback.print_exc()
# Still try to save whatever we have
merged_dir = OUTPUT_DIR / "grpo_adapter"
return merged_dir
def step_5_push_to_hub(model_dir: Path):
"""Push fine-tuned model to HF Hub."""
log.info("=" * 60)
log.info("STEP 5: Push to HuggingFace Hub")
log.info("=" * 60)
if not HF_TOKEN:
log.warning("No HF_TOKEN — skipping push. Model saved at %s", model_dir)
return
if not model_dir.exists():
log.error("Model dir %s does not exist — nothing to push", model_dir)
return
from huggingface_hub import HfApi
api = HfApi(token=HF_TOKEN)
# Create model repo
api.create_repo(repo_id=MODEL_REPO, exist_ok=True, private=False)
# Upload all files
api.upload_folder(
folder_path=str(model_dir),
repo_id=MODEL_REPO,
commit_message="FATHOM GRPO fine-tuned model",
)
log.info("Model pushed to https://huggingface.co/%s", MODEL_REPO)
def main():
log.info("=" * 60)
log.info("FATHOM Training Orchestrator")
log.info("=" * 60)
log.info("Config:")
log.info(" Model: %s", "0.5B smoke" if USE_SMOKE_MODEL else "1.5B full")
log.info(" Env URL: %s", ENV_SPACE_URL)
log.info(" Output: %s", OUTPUT_DIR)
log.info(" Push to: %s", MODEL_REPO if HF_TOKEN else "(no token)")
overall_start = time.time()
try:
step_1_verify_env()
step_2_generate_dataset()
adapter_dir = step_3_sft_warmstart()
merged_dir = step_4_grpo_training()
step_5_push_to_hub(merged_dir)
except Exception as e:
log.error("FATAL: %s", e)
import traceback
traceback.print_exc()
sys.exit(1)
total_min = (time.time() - overall_start) / 60
log.info("=" * 60)
log.info("TRAINING COMPLETE in %.1f minutes", total_min)
log.info("=" * 60)
# Keep container alive so logs are readable
log.info("Container will stay alive for 10 min for log inspection...")
time.sleep(600)
if __name__ == "__main__":
main()