AriaLM / src /s05_trainer.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
8.57 kB
"""
Training loop with:
- Mixed precision (AMP) for memory savings
- Gradient checkpointing for OOM prevention
- Cosine annealing with warmup
- Early stopping
- Gradient accumulation for larger effective batch
- TensorBoard logging
"""
import logging
import math
import time
from pathlib import Path
from typing import Optional
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter
from src.s01_config import TrainConfig, PathConfig, get_device
from src.s04_model import MusicTransformer
logger = logging.getLogger(__name__)
class CosineWarmupScheduler:
"""Cosine annealing LR with linear warmup."""
def __init__(self, optimizer, warmup_steps: int, total_steps: int, min_lr: float = 1e-6):
self.optimizer = optimizer
self.warmup_steps = warmup_steps
self.total_steps = total_steps
self.min_lr = min_lr
self.base_lrs = [pg["lr"] for pg in optimizer.param_groups]
self.step_count = 0
def step(self):
self.step_count += 1
for pg, base_lr in zip(self.optimizer.param_groups, self.base_lrs):
if self.step_count < self.warmup_steps:
lr = base_lr * self.step_count / max(1, self.warmup_steps)
else:
progress = (self.step_count - self.warmup_steps) / max(
1, self.total_steps - self.warmup_steps
)
lr = self.min_lr + (base_lr - self.min_lr) * 0.5 * (1 + math.cos(math.pi * progress))
pg["lr"] = lr
def get_lr(self) -> float:
return self.optimizer.param_groups[0]["lr"]
class Trainer:
"""Handles the full training pipeline with memory-efficient techniques."""
def __init__(
self,
model: MusicTransformer,
train_loader: DataLoader,
val_loader: DataLoader,
train_config: TrainConfig,
path_config: PathConfig,
):
self.model = model
self.train_loader = train_loader
self.val_loader = val_loader
self.config = train_config
self.paths = path_config
self.device = get_device()
# Enable gradient checkpointing
if train_config.grad_checkpoint:
self.model.grad_checkpoint = True
logger.info("Gradient checkpointing ENABLED")
self.model.to(self.device)
# Optimizer: AdamW with weight decay (decoupled)
self.optimizer = torch.optim.AdamW(
self.model.parameters(),
lr=train_config.learning_rate,
weight_decay=train_config.weight_decay,
betas=(0.9, 0.95),
fused=torch.cuda.is_available(), # Fused optimizer on CUDA
)
# LR scheduler
total_steps = len(train_loader) * train_config.max_epochs // train_config.grad_accum_steps
self.scheduler = CosineWarmupScheduler(
self.optimizer, train_config.warmup_steps, total_steps
)
# Mixed precision scaler
use_amp = train_config.use_amp and torch.cuda.is_available()
self.scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
if use_amp:
self.amp_dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
else:
self.amp_dtype = torch.float32
self.use_amp = use_amp
# TensorBoard
self.writer = SummaryWriter(log_dir=str(path_config.log_dir))
# Tracking
self.global_step = 0
self.best_val_loss = float("inf")
self.patience_counter = 0
def train(self):
"""Main training loop."""
logger.info(f"Starting training on {self.device}")
logger.info(f"Model params: {self.model.count_parameters():,}")
logger.info(f"Train batches: {len(self.train_loader)}, Val batches: {len(self.val_loader)}")
for epoch in range(1, self.config.max_epochs + 1):
t0 = time.time()
train_loss = self._train_epoch(epoch)
val_loss = self._validate()
elapsed = time.time() - t0
logger.info(
f"Epoch {epoch}/{self.config.max_epochs} | "
f"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | "
f"LR: {self.scheduler.get_lr():.2e} | Time: {elapsed:.1f}s"
)
self.writer.add_scalars("loss", {"train": train_loss, "val": val_loss}, epoch)
self.writer.add_scalar("lr", self.scheduler.get_lr(), epoch)
# Early stopping check
if val_loss < self.best_val_loss - self.config.min_delta:
self.best_val_loss = val_loss
self.patience_counter = 0
self._save_checkpoint("best.pt", epoch, val_loss)
logger.info(f" New best model saved (val_loss={val_loss:.4f})")
else:
self.patience_counter += 1
if self.patience_counter >= self.config.patience:
logger.info(f"Early stopping at epoch {epoch} (patience={self.config.patience})")
break
# Periodic checkpoint
if epoch % 5 == 0:
self._save_checkpoint(f"epoch_{epoch}.pt", epoch, val_loss)
self.writer.close()
logger.info("Training complete!")
def _train_epoch(self, epoch: int) -> float:
self.model.train()
total_loss = 0.0
n_batches = 0
self.optimizer.zero_grad(set_to_none=True)
for batch_idx, (input_ids, targets) in enumerate(self.train_loader):
input_ids = input_ids.to(self.device, non_blocking=True)
targets = targets.to(self.device, non_blocking=True)
# Mixed precision forward
with torch.amp.autocast(
device_type=self.device.type,
dtype=self.amp_dtype,
enabled=self.use_amp,
):
_, loss = self.model(input_ids, targets)
loss = loss / self.config.grad_accum_steps
# Backward with gradient scaling
self.scaler.scale(loss).backward()
if (batch_idx + 1) % self.config.grad_accum_steps == 0:
self.scaler.unscale_(self.optimizer)
nn.utils.clip_grad_norm_(self.model.parameters(), self.config.max_grad_norm)
self.scaler.step(self.optimizer)
self.scaler.update()
self.optimizer.zero_grad(set_to_none=True)
self.scheduler.step()
self.global_step += 1
total_loss += loss.item() * self.config.grad_accum_steps
n_batches += 1
if (batch_idx + 1) % self.config.log_interval == 0:
avg = total_loss / n_batches
logger.info(
f" Epoch {epoch} [{batch_idx+1}/{len(self.train_loader)}] "
f"loss={avg:.4f} lr={self.scheduler.get_lr():.2e}"
)
return total_loss / max(1, n_batches)
@torch.no_grad()
def _validate(self) -> float:
self.model.eval()
total_loss = 0.0
n_batches = 0
for input_ids, targets in self.val_loader:
input_ids = input_ids.to(self.device, non_blocking=True)
targets = targets.to(self.device, non_blocking=True)
with torch.amp.autocast(
device_type=self.device.type,
dtype=self.amp_dtype,
enabled=self.use_amp,
):
_, loss = self.model(input_ids, targets)
total_loss += loss.item()
n_batches += 1
return total_loss / max(1, n_batches)
def _save_checkpoint(self, name: str, epoch: int, val_loss: float):
path = self.paths.checkpoint_dir / name
torch.save(
{
"epoch": epoch,
"model_state_dict": self.model.state_dict(),
"optimizer_state_dict": self.optimizer.state_dict(),
"val_loss": val_loss,
"global_step": self.global_step,
"config": self.model.config,
},
path,
)
def load_checkpoint(self, path: Path):
ckpt = torch.load(path, map_location=self.device, weights_only=False)
self.model.load_state_dict(ckpt["model_state_dict"])
self.optimizer.load_state_dict(ckpt["optimizer_state_dict"])
self.global_step = ckpt.get("global_step", 0)
logger.info(f"Loaded checkpoint: {path} (epoch={ckpt['epoch']}, val_loss={ckpt['val_loss']:.4f})")