File size: 7,746 Bytes
8787bd3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
"""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()