fugu-lite / src /fugu_lite /train_sft.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
3.64 kB
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