"""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"]