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