fathom-code / train /sft.py
23f2002275
fix(sft): TRL 1.2 SFTConfig renamed max_seq_length -> max_length
483642a
Raw
History Blame Contribute Delete
4.34 kB
"""FATHOM SFT warm-start — TRN-02.
Wraps TRL SFTTrainer with a chat-template-aware formatting_func and the
STACK §6 safe-save sequence (adapter-only first, then optional HF Hub push).
Single source of truth for chat formatting: tokenizer.apply_chat_template.
DO NOT hand-concatenate <|im_start|>/<|im_end|> strings (STACK §3.3).
"""
from __future__ import annotations
import logging
import os
from pathlib import Path
from typing import Any
from omegaconf import DictConfig
log = logging.getLogger("fathom.train.sft")
def run_sft(
cfg: DictConfig,
model: Any,
tokenizer: Any,
dataset: Any | None = None,
) -> Path:
"""Build TRL SFTTrainer, train on cfg.data.sft_traces_path, save adapter.
Args:
cfg: composed Hydra DictConfig (reads cfg.train.*, cfg.data.sft_traces_path,
cfg.output_dir, cfg.hub.push, cfg.hub.repo_id).
model: Unsloth-patched PeftModel (from train/model_load.py).
tokenizer: HF tokenizer with chat_template set.
dataset: optional pre-loaded HF datasets.Dataset; if None, load JSONL
from cfg.data.sft_traces_path.
Returns:
Path to saved adapter dir: `{cfg.output_dir}/sft_adapter/`.
"""
# Lazy imports so `import train.sft` is cheap for unit tests
from trl import SFTConfig, SFTTrainer # type: ignore # noqa: F401
from datasets import load_dataset # type: ignore # noqa: F401
# TRN-02: single source of truth for chat formatting (STACK §3.3)
def _formatting_func(example: dict) -> str:
return tokenizer.apply_chat_template(
example["messages"],
tokenize=False,
add_generation_prompt=False,
)
# Dataset loading
if dataset is None:
sft_path = Path(str(cfg.data.sft_traces_path))
if not sft_path.exists():
raise FileNotFoundError(
f"TRN-02: sft_traces_path not found at {sft_path} — run DATA-06 first"
)
dataset = load_dataset("json", data_files=str(sft_path), split="train")
log.info(
"TRN-02 SFT starting: dataset_rows=%d max_seq_length=%d lr=%s",
len(dataset),
cfg.train.max_seq_length,
cfg.train.learning_rate,
)
# Log first example preview (STACK §3.3 diff-test aid)
if len(dataset) > 0:
preview = _formatting_func(dataset[0])
log.info("TRN-02 first tokenized example (first 200 chars): %s", preview[:200])
# Build SFTConfig from Hydra cfg — type-cast every value (OmegaConf safety)
sft_config = SFTConfig(
output_dir=str(Path(str(cfg.output_dir)) / "sft_run"),
learning_rate=float(cfg.train.learning_rate),
num_train_epochs=float(cfg.train.num_train_epochs),
max_length=int(cfg.train.max_seq_length),
per_device_train_batch_size=int(cfg.train.per_device_train_batch_size),
gradient_accumulation_steps=int(cfg.train.gradient_accumulation_steps),
optim=str(cfg.train.optim),
bf16=bool(cfg.train.bf16),
save_strategy=str(cfg.train.save_strategy),
seed=int(cfg.seed),
report_to=[], # No implicit W&B in SFT; GRPO owns W&B (REW-03)
logging_steps=5,
packing=False,
)
trainer = SFTTrainer(
model=model,
processing_class=tokenizer,
args=sft_config,
train_dataset=dataset,
formatting_func=_formatting_func,
)
trainer.train()
# TRN-02: STACK §6 save sequence — adapter-only FIRST
adapter_dir = Path(str(cfg.output_dir)) / "sft_adapter"
model.save_pretrained(str(adapter_dir)) # TRN-02 adapter-only save
tokenizer.save_pretrained(str(adapter_dir))
log.info("TRN-02 adapter saved to %s", adapter_dir)
# Optional HF Hub push — gated on both cfg flag AND env var
hub_push = getattr(cfg, "hub", None)
if hub_push is not None and bool(getattr(hub_push, "push", False)):
token = os.environ.get("HF_TOKEN")
if token:
repo_id = str(cfg.hub.repo_id)
model.push_to_hub(repo_id, token=token)
tokenizer.push_to_hub(repo_id, token=token)
log.info("TRN-02 adapter pushed to %s", repo_id)
else:
log.info("TRN-02 hub.push=true but HF_TOKEN not set — skipping push")
return adapter_dir
__all__ = ["run_sft"]