| """FATHOM GRPO training scaffold — TRN-03. |
| |
| Contract: |
| - Consumes (model, tokenizer) from train.model_load.load_model_and_tokenizer |
| - Loads SFT adapter from {cfg.output_dir}/sft_adapter/ if present (Plan 02 output) |
| - Builds trl.GRPOTrainer with vllm_mode='colocate' (STACK §4 + §10.4 — 'server' |
| mode breaks multi-turn OpenEnv rollouts per TRL #4543) |
| - Saves adapter-only FIRST, then merged_16bit (STACK §6 + §10.1 — |
| NEVER merged_4bit / merged_4bit_forced: corrupt under QLoRA) |
| |
| REW-03: per-component W&B logging is wired via the instrumented reward_fn |
| callback pattern (see make_reward_fn in rewards/compose.py). |
| """ |
| from __future__ import annotations |
|
|
| import logging |
| import os |
| import inspect |
| import statistics |
| from pathlib import Path |
| from typing import Any, Callable |
|
|
| from omegaconf import DictConfig |
|
|
| log = logging.getLogger("fathom.train.grpo") |
|
|
|
|
| def run_grpo( |
| cfg: DictConfig, |
| model: Any, |
| tokenizer: Any, |
| reward_fn: Callable, |
| env_url: str, |
| ) -> Path: |
| """Load SFT adapter, build GRPOTrainer, train, save merged_16bit. |
| |
| Args: |
| cfg: composed Hydra DictConfig (reads cfg.train.*, cfg.output_dir, cfg.seed, cfg.hub.*). |
| model: Unsloth-patched PeftModel returned by load_model_and_tokenizer(cfg). |
| tokenizer: HF tokenizer. |
| reward_fn: callable(prompts, completions, **kwargs) -> list[float] |
| (TRL's reward-function contract). |
| env_url: OpenEnv HTTP URL (http://localhost:8001 locally, HF Space URL at venue). |
| |
| Returns: |
| Path to saved merged model dir: `{cfg.output_dir}/grpo_merged_16bit/`. |
| """ |
| |
| from trl import GRPOTrainer, GRPOConfig |
| import peft |
| from datasets import load_dataset |
|
|
| |
| |
| |
| |
| |
| |
| sft_adapter_dir = Path(str(cfg.output_dir)) / "sft_adapter" |
| if sft_adapter_dir.exists(): |
| try: |
| if isinstance(model, peft.PeftModel): |
| base = model.unload() |
| else: |
| base = model |
| model = peft.PeftModel.from_pretrained( |
| base, str(sft_adapter_dir), is_trainable=True |
| ) |
| log.info("TRN-03 SFT adapter loaded from %s (after unloading empty wrap)", sft_adapter_dir) |
| except Exception as e: |
| log.warning("TRN-03 SFT adapter load failed (%s) — continuing with base LoRA", e) |
| else: |
| log.info( |
| "TRN-03: no SFT adapter at %s — proceeding with base LoRA (smoke path only)", |
| sft_adapter_dir, |
| ) |
|
|
| |
| |
| |
| use_vllm_env = os.environ.get("FATHOM_USE_VLLM", "1").strip() |
| use_vllm = use_vllm_env not in ("0", "false", "False", "") |
| if use_vllm: |
| assert str(cfg.train.vllm_mode) == "colocate", ( |
| "TRN-03 gate: vllm_mode must be 'colocate' for multi-turn OpenEnv (STACK §10.4)" |
| ) |
|
|
| |
| |
| |
| cfg_sig = inspect.signature(GRPOConfig) |
| params = cfg_sig.parameters |
| supported = set(params.keys()) |
| supports_kwargs = any( |
| p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values() |
| ) |
| kwargs: dict[str, Any] = {} |
|
|
| def _set_if_supported(candidates: list[str], value: Any) -> None: |
| for key in candidates: |
| if supports_kwargs or key in supported: |
| kwargs[key] = value |
| return |
|
|
| _set_if_supported(["output_dir"], str(Path(str(cfg.output_dir)) / "grpo_run")) |
| _set_if_supported(["num_generations"], int(cfg.train.num_generations)) |
| _set_if_supported(["beta"], float(cfg.train.beta)) |
| _set_if_supported(["learning_rate"], float(cfg.train.learning_rate)) |
| _set_if_supported(["max_grad_norm"], float(cfg.train.max_grad_norm)) |
| _set_if_supported(["bf16"], bool(cfg.train.bf16)) |
| _set_if_supported( |
| ["max_prompt_length", "prompt_max_length", "max_prompt_tokens"], |
| int(cfg.train.max_prompt_length), |
| ) |
| _set_if_supported( |
| ["max_completion_length", "completion_max_length", "max_new_tokens"], |
| int(cfg.train.max_completion_length), |
| ) |
| _set_if_supported(["optim"], str(cfg.train.optim)) |
| _set_if_supported(["max_steps"], int(cfg.train.max_steps)) |
| _set_if_supported(["save_steps"], int(cfg.train.save_steps)) |
| _set_if_supported(["seed"], int(cfg.seed)) |
| _set_if_supported(["use_vllm"], use_vllm) |
| if use_vllm: |
| _set_if_supported(["vllm_mode"], str(cfg.train.vllm_mode)) |
| _set_if_supported( |
| ["vllm_gpu_memory_utilization"], |
| float(cfg.train.vllm_gpu_memory_utilization), |
| ) |
| log.info("TRN-03 vLLM rollout enabled (vllm_mode=%s)", cfg.train.vllm_mode) |
| else: |
| log.warning( |
| "TRN-03 FATHOM_USE_VLLM=0 — using HF generate() for " |
| "rollouts. Slower but importance_sampling_ratio stays " |
| "near 1.0 (avoids QLoRA + vLLM merge drift bug)." |
| ) |
| wb_key = os.environ.get("WANDB_API_KEY", "").strip() |
| |
| use_wandb = len(wb_key) >= 40 and all(c.isalnum() or c == "_" for c in wb_key) |
| _set_if_supported(["report_to"], ["wandb"] if use_wandb else []) |
| if not use_wandb: |
| log.warning( |
| "TRN-03 W&B disabled: WANDB_API_KEY missing or wrong length " |
| "(have %d chars, need 40). Training will continue without " |
| "W&B logging; metrics still printed to stdout.", len(wb_key) |
| ) |
| _set_if_supported(["logging_steps"], 1) |
| |
| |
| _set_if_supported(["remove_unused_columns"], False) |
| |
| _set_if_supported(["per_device_train_batch_size"], int(cfg.train.num_generations)) |
| _set_if_supported(["gradient_checkpointing"], True) |
|
|
| grpo_config = GRPOConfig(**kwargs) |
|
|
| |
| try: |
| import wandb |
| except ImportError: |
| wandb = None |
|
|
| def _instrumented_reward_fn(prompts, completions, **kwargs): |
| rewards = reward_fn(prompts, completions, **kwargs) |
| try: |
| if wandb is not None and wandb.run is not None: |
| log_dict = { |
| "reward/composite_mean": sum(rewards) / max(len(rewards), 1), |
| "reward/composite_std": ( |
| statistics.stdev(rewards) if len(rewards) > 1 else 0.0 |
| ), |
| } |
| cm = kwargs.get("_component_means", {}) |
| for k, v in cm.items(): |
| log_dict[f"reward/{k}_mean"] = float(v) |
| wandb.log(log_dict) |
| except Exception: |
| pass |
| return rewards |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| _ = env_url |
|
|
| train_path = Path(str(cfg.data.train_path)) |
| if not train_path.exists(): |
| raise FileNotFoundError( |
| f"TRN-03: train_path not found at {train_path} — run DATA-* first" |
| ) |
| raw_ds = load_dataset("json", data_files=str(train_path), split="train") |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| sys_msg = "You are FATHOM, a recursive language model with a Python REPL sandbox. You can read a long document via the variable `ctx` and call `llm(prompt, chunk)` for sub-queries. Think step by step. Emit your final answer inside <answer>...</answer>." |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| max_prompt_tok = int(cfg.train.max_prompt_length) |
| chat_overhead_tok = 256 |
| ctx_budget_tok = max(256, max_prompt_tok - chat_overhead_tok) |
|
|
| def _truncate_to_tokens(text: str, max_tokens: int) -> str: |
| if not text: |
| return "" |
| ids = tokenizer.encode(text, add_special_tokens=False) |
| if len(ids) <= max_tokens: |
| return text |
| |
| |
| return tokenizer.decode(ids[-max_tokens:], skip_special_tokens=True) |
|
|
| def _to_prompt(example: dict) -> dict: |
| ctx_full = example.get("context", "") or "" |
| ctx_truncated = _truncate_to_tokens(ctx_full, ctx_budget_tok) |
| |
| user_content = ( |
| f"{example.get('prompt', '')}\n\n" |
| f"[Document excerpt]:\n{ctx_truncated}" |
| ) |
| msgs = [ |
| {"role": "system", "content": sys_msg}, |
| {"role": "user", "content": user_content}, |
| ] |
| prompt_str = tokenizer.apply_chat_template( |
| msgs, tokenize=False, add_generation_prompt=True |
| ) |
| return { |
| "prompt": prompt_str, |
| "gold_answer": str(example.get("gold_answer", "")), |
| "prompt_token_count": int(example.get("context_length", 0)) // 4, |
| } |
|
|
| train_dataset = raw_ds.map( |
| _to_prompt, |
| remove_columns=[c for c in raw_ds.column_names if c not in {"prompt"}], |
| load_from_cache_file=False, |
| ) |
| log.info("TRN-03 train_dataset ready: rows=%d", len(train_dataset)) |
|
|
| |
| if "remove_unused_columns" in inspect.signature(GRPOConfig).parameters: |
| grpo_config.remove_unused_columns = False |
|
|
| trainer = GRPOTrainer( |
| model=model, |
| processing_class=tokenizer, |
| args=grpo_config, |
| reward_funcs=[_instrumented_reward_fn], |
| train_dataset=train_dataset, |
| ) |
| log.info("TRN-03 GRPOTrainer constructed (no env tools — pure prompt→completion→reward)") |
|
|
| |
| |
| |
| import json as _json |
| _sft = _json.loads(open(str(cfg.data.sft_traces_path) if hasattr(cfg.data, "sft_traces_path") else "data/sft_traces.jsonl", encoding="utf-8").readline()) |
| _sft_prefix = tokenizer.apply_chat_template(_sft["messages"][:2], tokenize=False, add_generation_prompt=True)[:200] |
| _grpo_first = train_dataset[0]["prompt"][:200] |
| assert _sft_prefix.split("Question:")[0] == _grpo_first.split("Question:")[0], ( |
| "SFT/GRPO chat-template prefix drift detected — see ANTIGRAVITY_BRIEF.md §A.3.1" |
| ) |
|
|
| trainer.train() |
|
|
| |
| adapter_dir = Path(str(cfg.output_dir)) / "grpo_adapter" |
| model.save_pretrained(str(adapter_dir)) |
| tokenizer.save_pretrained(str(adapter_dir)) |
| log.info("TRN-03 GRPO adapter saved to %s", adapter_dir) |
|
|
| |
| hub_cfg = getattr(cfg, "hub", None) |
| if hub_cfg is not None and bool(getattr(hub_cfg, "push", False)): |
| token = os.environ.get("HF_TOKEN") |
| if token: |
| repo_id = str(cfg.hub.repo_id) + "-adapter" |
| model.push_to_hub(repo_id, token=token) |
| log.info("TRN-03 GRPO adapter pushed to %s", repo_id) |
|
|
| |
| |
| save_method = "merged_16bit" |
| assert save_method == "merged_16bit", ( |
| "STACK §10.1 anti-pattern: merged_4bit / merged_4bit_forced are corrupt under QLoRA" |
| ) |
| merged_dir = Path(str(cfg.output_dir)) / "grpo_merged_16bit" |
| try: |
| model.save_pretrained_merged(str(merged_dir), tokenizer, save_method=save_method) |
| log.info("TRN-03 merged_16bit saved to %s", merged_dir) |
| return merged_dir |
| except Exception as e: |
| log.warning("TRN-03 save_pretrained_merged failed (%s) — falling back to peft merge", e) |
| try: |
| merged_model = model.merge_and_unload() |
| merged_model.save_pretrained(str(merged_dir)) |
| tokenizer.save_pretrained(str(merged_dir)) |
| log.info("TRN-03 fallback peft merge saved to %s", merged_dir) |
| return merged_dir |
| except Exception as e2: |
| log.error("TRN-03 peft fallback also failed (%s) — returning adapter_dir", e2) |
| return adapter_dir |
|
|
|
|
| __all__ = ["run_grpo"] |
|
|