"""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/`.
"""
# Lazy imports so `import train.grpo` is cheap for unit tests
from trl import GRPOTrainer, GRPOConfig # type: ignore # noqa: F401
import peft # type: ignore # noqa: F401
from datasets import load_dataset # type: ignore # noqa: F401
# TRN-03 step 1: Load SFT adapter if present (Plan 02 output).
# NOTE: `load_model_and_tokenizer` already wraps the base model with a
# fresh LoRA via PEFT. Calling `PeftModel.from_pretrained(model, …)` again
# double-wraps it — visible in the warning "Found missing adapter keys"
# with paths like `base_model.model.base_model.model.…`. The fix is to
# unload the empty LoRA first, then attach the SFT adapter to the base.
sft_adapter_dir = Path(str(cfg.output_dir)) / "sft_adapter"
if sft_adapter_dir.exists():
try:
if isinstance(model, peft.PeftModel):
base = model.unload() # strip empty LoRA wrap
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,
)
# TRN-03 step 2: FATHOM_USE_VLLM env var gates vLLM rollout.
# Set FATHOM_USE_VLLM=0 to use HF generate() for rollouts (slower but
# avoids IS-ratio collapse under QLoRA + vLLM merge drift).
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)"
)
# TRN-03 step 2: Build GRPOConfig from cfg.train.
# TRL minor versions have changed some GRPOConfig field names; select only
# kwargs that exist in the installed signature and map common aliases.
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()
# Accept both old-style (40-char alnum) and new wandb_v1_ keys (contain underscores).
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)
# Keep TRL from dropping our extra columns (gold_answer etc.) so they
# reach the reward function as kwargs.
_set_if_supported(["remove_unused_columns"], False)
# Per-device batch needs to be divisible by num_generations (TRL gate).
_set_if_supported(["per_device_train_batch_size"], int(cfg.train.num_generations))
_set_if_supported(["gradient_checkpointing"], True)
grpo_config = GRPOConfig(**kwargs)
# REW-03: wrap reward_fn to log per-component scalars to W&B
try:
import wandb # type: ignore
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
# TRN-03 step 3: Build train_dataset from data/train.jsonl.
#
# Why no `env=` / `environment_url=` / `environment_factory` kwargs?
# - TRL 1.2.0's GRPOTrainer.__init__ only accepts env interaction via
# `tools=` (needs transformers>=5.0), `environment_factory=`
# (needs transformers>=5.2), or `rollout_func=`. We're on
# transformers==4.56.2, so the first two raise. The third requires a
# custom multi-turn rollout implementation we don't have time to
# harden.
# - Our reward function (rewards.compose) operates on (prompt, completion,
# gold_answer, prompt_token_count, llm_call_count) — zero env
# interaction needed. The OpenEnv server stays up as the deployable
# artifact (judges hit /healthz, the demo Streamlit drives recursion at
# inference time).
# The reward function ignores env_url at training time but we keep the
# signature so smoke / unit tests don't break.
_ = 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")
# Map to TRL-friendly columns. `prompt` is the chat-templated string;
# `gold_answer`, `prompt_token_count`, `llm_call_count` come along as
# extra columns and TRL forwards them to the reward function as kwargs
# (because `remove_unused_columns=False` below).
# IMPORTANT: must match the system message used in SFT traces
# (data/sft_traces.jsonl). REW-01 format_gate is a multiplicative gate —
# if the completion lacks ..., composite reward = 0.0.
# Previous wording ("shortest exact answer span") never told the model
# about the tag, so 50/50 GRPO steps had reward=0.0 (job 69ece94a).
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 ...."
# vLLM hard-checks final prompt length against the model's max_position_embeddings
# (32 768 for Qwen 2.5 Coder 0.5B/1.5B). Our train.jsonl `context` field is the
# full long document (often >40K tokens). `max_prompt_length` in GRPOConfig is a
# post-tokenization cap that TRL applies AFTER vLLM has already rejected the
# request (`vllm.exceptions.VLLMValidationError: prompt contains 37220 tokens`).
# We must truncate inside `_to_prompt`, BEFORE chat templating.
#
# Budget: keep the prompt comfortably under cfg.train.max_prompt_length so the
# system message + chat-template overhead don't push us over. Tail-truncate the
# context (recent text usually contains the answer span in our synthetic data).
max_prompt_tok = int(cfg.train.max_prompt_length)
chat_overhead_tok = 256 # system msg + chat template wrappers + question
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
# Keep the tail: synthetic gold answers are sampled across the doc, so tail
# is no worse than head and is cheaper to slice.
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)
# CRITICAL: must match data/sft_traces.jsonl user-message shape exactly.
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))
# Tell TRL to keep the extra columns so they appear in the reward fn kwargs.
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)")
# TRN-03 step 4: Train
# Pre-flight: tokenize one example and confirm the chat-template prefix is
# the byte-identical match of an SFT trace prefix.
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()
# TRN-03 step 5: STACK §6 save sequence — adapter-only FIRST
adapter_dir = Path(str(cfg.output_dir)) / "grpo_adapter"
model.save_pretrained(str(adapter_dir)) # adapter-only first (R4 ruin-mode insurance)
tokenizer.save_pretrained(str(adapter_dir))
log.info("TRN-03 GRPO adapter saved to %s", adapter_dir)
# Optional Hub push for adapter (checkpoint insurance every save_steps)
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)
# TRN-03 step 5c: Merged save — HARD ASSERT: only merged_16bit is safe (STACK §10.1)
# NEVER merged_4bit / merged_4bit_forced: corrupt under QLoRA (issues #1267 #2339 #1791)
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"]