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