Nova / src /trainer /train_loop.py
kings1's picture
Upload folder using huggingface_hub
23ea6bd verified
Raw History Blame Contribute Delete
5.48 kB
"""
Training Loop & Optimizer Engine for Nova 1.0 (Multi-Core Scaled)
Implements Deep Supervision & Adaptive Computation Time (ACT) loss routines.
Optimized for AMD Ryzen AI Max+ 395 (32 CPU threads) & AMD ROCm.
"""
import os
import math
from typing import Dict
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from tqdm import tqdm
from config.model_config import Nova1Config
from src.model.nova1_hrm import Nova1HRM
class Nova1Trainer:
def __init__(self, model: Nova1HRM, config: Nova1Config, dataloader: DataLoader):
self.model = model
self.config = config
self.dataloader = dataloader
self.device = torch.device(config.device)
self.model.to(self.device)
# Utilize all CPU threads on AMD Ryzen AI Max+ 395
if hasattr(torch, "set_num_threads"):
torch.set_num_threads(config.num_threads)
# CrossEntropyLoss (ignore padding tokens)
self.criterion = nn.CrossEntropyLoss(ignore_index=config.pad_token_id)
# AdamW Optimizer with weight decay
self.optimizer = torch.optim.AdamW(
self.model.parameters(),
lr=config.learning_rate,
weight_decay=config.weight_decay,
betas=(0.9, 0.95)
)
# Automatic Mixed Precision GradScaler (only needed for float16)
self.use_amp = (config.device == "cuda")
self.use_scaler = (config.device == "cuda" and config.dtype == "float16")
self.scaler = torch.amp.GradScaler("cuda", enabled=self.use_scaler)
def adjust_learning_rate(self, step: int, total_steps: int):
"""Cosine annealing learning rate schedule with linear warmup."""
if step < self.config.warmup_steps:
lr = self.config.learning_rate * (step / max(1, self.config.warmup_steps))
else:
progress = (step - self.config.warmup_steps) / max(1, total_steps - self.config.warmup_steps)
lr = self.config.min_learning_rate + 0.5 * (self.config.learning_rate - self.config.min_learning_rate) * (
1.0 + math.cos(math.pi * progress)
)
for param_group in self.optimizer.param_groups:
param_group["lr"] = lr
return lr
def train_epoch(self, epoch: int, total_epochs: int) -> float:
self.model.train()
total_loss = 0.0
steps_per_epoch = len(self.dataloader)
total_training_steps = steps_per_epoch * total_epochs
for step, (input_ids, target_ids) in enumerate(self.dataloader):
global_step = epoch * steps_per_epoch + step
lr = self.adjust_learning_rate(global_step, total_training_steps)
input_ids = input_ids.to(self.device)
target_ids = target_ids.to(self.device)
self.optimizer.zero_grad()
step_loss = 0.0
states = None
# Deep Supervision Loop over M_max segments
for m in range(self.config.max_segments):
with torch.amp.autocast("cuda", enabled=self.use_amp, dtype=self.config.get_torch_dtype()):
logits, (z_H, z_L), q_values = self.model(input_ids, states=states)
# Language Modeling Cross Entropy Loss
loss_lm = self.criterion(logits.view(-1, self.config.vocab_size), target_ids.view(-1))
segment_loss = loss_lm
# Backpropagate segment loss (Equilibrium 1-Step gradient)
if self.use_scaler:
self.scaler.scale(segment_loss).backward()
else:
segment_loss.backward()
step_loss += segment_loss.item()
# Detach hidden states before passing into next supervision segment (Paper Section 2: Deep supervision)
states = (z_H.detach(), z_L.detach())
# Gradient Clipping & Optimizer step
if self.use_scaler:
self.scaler.unscale_(self.optimizer)
torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.max_grad_norm)
self.scaler.step(self.optimizer)
self.scaler.update()
else:
torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.max_grad_norm)
self.optimizer.step()
avg_step_loss = step_loss / self.config.max_segments
total_loss += avg_step_loss
if (step + 1) % 20 == 0 or (step + 1) == steps_per_epoch:
print(f"Epoch {epoch+1}/{total_epochs} [{step+1}/{steps_per_epoch}] — Step Loss: {avg_step_loss:.4f} | LR: {lr:.6f}", flush=True)
return total_loss / steps_per_epoch
def save_checkpoint(self, checkpoint_path: str):
os.makedirs(os.path.dirname(checkpoint_path), exist_ok=True)
checkpoint = {
"model_state": self.model.state_dict(),
"optimizer_state": self.optimizer.state_dict(),
"config": self.config
}
torch.save(checkpoint, checkpoint_path)
print(f"Checkpoint saved to {checkpoint_path}")
def load_checkpoint(self, checkpoint_path: str):
checkpoint = torch.load(checkpoint_path, map_location=self.device)
self.model.load_state_dict(checkpoint["model_state"])
self.optimizer.load_state_dict(checkpoint["optimizer_state"])
print(f"Loaded checkpoint from {checkpoint_path}")