Download src/trainer/train_loop.py from kings1/Nova: direct link, hf CLI and curl.
- Browser
- Download file 5.48 kB
-
https://huggingface.co/kings1/Nova/resolve/main/src/trainer/train_loop.py
- Command line
-
hf download hf://kings1/Nova/src/trainer/train_loop.py
-
curl -L -o train_loop.py https://huggingface.co/kings1/Nova/resolve/main/src/trainer/train_loop.py
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}") | |