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