""" 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}")