File size: 3,643 Bytes
88e15cd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from pathlib import Path

import torch
import torch.nn.functional as F
from rich.console import Console

from .config import SFTConfig
from .io import write_json
from .schemas import RewardRecord
from .training_common import (
    evaluate_loader,
    make_loader,
    make_model,
    move_batch,
    optimizer_parameters,
    seed_everything,
)

console = Console()


def soft_label_loss(logits: torch.Tensor, rewards: torch.Tensor, temperature: float) -> torch.Tensor:
    if temperature <= 0:
        raise ValueError("soft_target_temperature must be positive")
    targets = F.softmax(rewards / temperature, dim=-1)
    return -(targets * F.log_softmax(logits, dim=-1)).sum(dim=-1).mean()


def train_sft(
    records: list[RewardRecord],
    config: SFTConfig,
    output_dir: str | Path,
    checkpoint: str | None = None,
) -> dict:
    seed_everything(config.training.seed)
    worker_ids = records[0].worker_ids
    model = make_model(config.model, worker_ids, checkpoint)
    train_loader = make_loader(
        records, model, "train", config.training.batch_size, shuffle=True
    )
    try:
        validation_loader = make_loader(
            records, model, "validation", config.training.batch_size, shuffle=False
        )
    except ValueError:
        validation_loader = None

    trainable, total = model.trainable_parameter_counts()
    console.print(
        f"Device={model.device_ref}; trainable={trainable:,}/{total:,} "
        f"({100 * trainable / total:.4f}%)"
    )
    optimizer = torch.optim.AdamW(
        optimizer_parameters(model),
        lr=config.training.learning_rate,
        weight_decay=config.training.weight_decay,
    )
    history: list[dict] = []
    global_step = 0
    optimizer.zero_grad(set_to_none=True)
    for epoch in range(config.training.epochs):
        model.train()
        epoch_losses = []
        for step, batch in enumerate(train_loader):
            inputs, rewards = move_batch(batch, model.device_ref)
            logits = model(**inputs)
            loss = soft_label_loss(logits, rewards, config.training.soft_target_temperature)
            (loss / config.training.gradient_accumulation_steps).backward()
            epoch_losses.append(float(loss.detach()))
            should_step = (
                (step + 1) % config.training.gradient_accumulation_steps == 0
                or step + 1 == len(train_loader)
            )
            if should_step:
                torch.nn.utils.clip_grad_norm_(
                    list(optimizer_parameters(model)), config.training.max_grad_norm
                )
                optimizer.step()
                optimizer.zero_grad(set_to_none=True)
                global_step += 1
                if global_step % config.training.log_every == 0:
                    console.print(f"SFT step={global_step} loss={epoch_losses[-1]:.4f}")
        epoch_report = {
            "epoch": epoch + 1,
            "train_loss": sum(epoch_losses) / max(1, len(epoch_losses)),
        }
        if validation_loader is not None:
            epoch_report["validation"] = evaluate_loader(model, validation_loader)
        history.append(epoch_report)
        console.print(epoch_report)

    output = Path(output_dir)
    model.save_checkpoint(
        output,
        metadata={"stage": "sft", "global_step": global_step, "history": history},
    )
    report = {
        "stage": "sft",
        "checkpoint": str(output),
        "worker_ids": worker_ids,
        "global_step": global_step,
        "history": history,
    }
    write_json(output / "training_report.json", report)
    return report