File size: 4,340 Bytes
071ba6b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
483642a
071ba6b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""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"]