| 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 |
|
|
|
|