| """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/`. |
| """ |
| |
| from trl import SFTConfig, SFTTrainer |
| from datasets import load_dataset |
|
|
| |
| def _formatting_func(example: dict) -> str: |
| return tokenizer.apply_chat_template( |
| example["messages"], |
| tokenize=False, |
| add_generation_prompt=False, |
| ) |
|
|
| |
| 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, |
| ) |
|
|
| |
| if len(dataset) > 0: |
| preview = _formatting_func(dataset[0]) |
| log.info("TRN-02 first tokenized example (first 200 chars): %s", preview[:200]) |
|
|
| |
| 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=[], |
| logging_steps=5, |
| packing=False, |
| ) |
|
|
| trainer = SFTTrainer( |
| model=model, |
| processing_class=tokenizer, |
| args=sft_config, |
| train_dataset=dataset, |
| formatting_func=_formatting_func, |
| ) |
| trainer.train() |
|
|
| |
| adapter_dir = Path(str(cfg.output_dir)) / "sft_adapter" |
| model.save_pretrained(str(adapter_dir)) |
| tokenizer.save_pretrained(str(adapter_dir)) |
| log.info("TRN-02 adapter saved to %s", adapter_dir) |
|
|
| |
| 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"] |
|
|