Premchan369/Q-TensorFormer
2185
1"""2Training utilities with budget-aware scheduling, energy metrics, and sweep support.3 4v3 features:5 - Budget-constrained training (auto-adjusts ranks to meet param/latency targets)6 - Energy estimation (FLOPs-based proxy)7 - Knowledge distillation support8 - Gradient monitoring and NaN detection9 - Checkpointing with metadata10"""11 12import torch13import torch.nn as nn14import torch.nn.functional as F15from torch.optim import AdamW16from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, LinearLR, SequentialLR17import math18import time19from typing import Optional, Dict, Tuple, List20from pathlib import Path21import json22 23from .config import ExperimentConfig24from .budget import BudgetTracker, EnergyEstimator25 26 27def create_optimizer(model: nn.Module, lr: float, weight_decay: float,28 betas: Tuple[float, float] = (0.9, 0.98),29 eps: float = 1e-8) -> AdamW:30 """Create AdamW optimizer with weight decay exclusion for norms/biases."""31 no_decay = ["bias", "LayerNorm.weight", "layernorm.weight", "ln.weight"]32 params = [33 {34 "params": [p for n, p in model.named_parameters()35 if p.requires_grad and not any(nd in n for nd in no_decay)],36 "weight_decay": weight_decay,37 },38 {39 "params": [p for n, p in model.named_parameters()40 if p.requires_grad and any(nd in n for nd in no_decay)],41 "weight_decay": 0.0,42 },43 ]44 return AdamW(params, lr=lr, betas=betas, eps=eps)45 46 47def create_scheduler(optimizer, warmup_steps: int, max_steps: int,48 lr_min_factor: float = 0.1, scheduler_type: str = "cosine"):49 """Create learning rate scheduler with warmup."""50 warmup = LinearLR(optimizer, start_factor=1e-3, end_factor=1.0,51 total_iters=warmup_steps)52 53 if scheduler_type == "cosine":54 main = CosineAnnealingWarmRestarts(55 optimizer, T_0=max_steps - warmup_steps,56 T_mult=1, eta_min=lr_min_factor * optimizer.param_groups[0]["lr"]57 )58 elif scheduler_type == "linear":59 main = LinearLR(optimizer, start_factor=1.0,60 end_factor=lr_min_factor,61 total_iters=max_steps - warmup_steps)62 else:63 main = LinearLR(optimizer, start_factor=1.0, end_factor=1.0,64 total_iters=max_steps - warmup_steps)65 66 return SequentialLR(optimizer, schedulers=[warmup, main],67 milestones=[warmup_steps])68 69 70def compute_perplexity(logits: torch.Tensor, targets: torch.Tensor,71 ignore_index: int = 0) -> float:72 """Compute perplexity with ignore_index."""73 loss = F.cross_entropy(74 logits.reshape(-1, logits.size(-1)),75 targets.reshape(-1),76 ignore_index=ignore_index,77 reduction="mean",78 )79 return math.exp(loss.item())80 81 82class Trainer:83 """84 Budget-aware Q-TensorFormer trainer.85 86 Tracks:87 - Perplexity (primary metric)88 - Model size (parameters)89 - Latency estimates90 - Energy consumption (FLOPs proxy)91 - Quantum call statistics92 - Rank adaptation trajectories93 """94 95 def __init__(self, model: nn.Module, config: ExperimentConfig,96 train_loader, val_loader=None, test_loader=None,97 device: str = "cpu", output_dir: str = None):98 self.model = model99 self.config = config100 self.train_loader = train_loader101 self.val_loader = val_loader102 self.test_loader = test_loader103 self.device = torch.device(device)104 self.output_dir = Path(output_dir or config.output_dir)105 106 self.model.to(self.device)107 108 total_steps = len(train_loader) * config.training.max_epochs109 self.optimizer = create_optimizer(110 model, config.training.learning_rate, config.training.weight_decay111 )112 self.scheduler = create_scheduler(113 self.optimizer,114 warmup_steps=config.training.warmup_steps,115 max_steps=total_steps,116 lr_min_factor=config.training.lr_min_factor,117 scheduler_type=config.training.lr_scheduler,118 )119 120 # Budget tracking121 self.budget_tracker = BudgetTracker(config.budget)122 self.energy_estimator = EnergyEstimator()123 124 # Logging125 self.metrics_history: List[Dict] = []126 self.grad_norms: List[float] = []127 128 def train_epoch(self, epoch: int) -> Dict:129 """Train for one epoch. Returns metrics dict."""130 self.model.train()131 self.model.reset_schedulers()132 total_loss = 0.0133 total_tokens = 0134 start_time = time.time()135 136 for step, (inputs, targets) in enumerate(self.train_loader):137 inputs, targets = inputs.to(self.device), targets.to(self.device)138 139 self.optimizer.zero_grad()140 141 logits, stats = self.model(inputs, return_stats=True)142 loss = F.cross_entropy(143 logits.reshape(-1, logits.size(-1)),144 targets.reshape(-1),145 ignore_index=0, # pad token146 )147 148 loss.backward()149 150 # Gradient monitoring151 grad_norm = torch.nn.utils.clip_grad_norm_(152 self.model.parameters(), self.config.training.max_grad_norm153 )154 self.grad_norms.append(grad_norm.item())155 156 # NaN check157 if torch.isnan(grad_norm) or torch.isinf(grad_norm):158 print(f"[WARN] NaN/Inf gradient at step {step}. Skipping update.")159 self.optimizer.zero_grad()160 continue161 162 self.optimizer.step()163 self.scheduler.step()164 165 total_loss += loss.item() * inputs.size(0) * inputs.size(1)166 total_tokens += inputs.size(0) * inputs.size(1)167 168 elapsed = time.time() - start_time169 avg_loss = total_loss / max(total_tokens, 1)170 ppl = math.exp(min(avg_loss, 20.0)) # Cap for stability171 172 # Budget metrics173 latency_est = self.budget_tracker.estimate_latency(174 self.model, self.config.model.max_seq_len175 )176 energy_est = self.energy_estimator.estimate(self.model)177 178 metrics = {179 "epoch": epoch,180 "train_loss": avg_loss,181 "train_ppl": ppl,182 "lr": self.optimizer.param_groups[0]["lr"],183 "grad_norm_mean": sum(self.grad_norms[-len(self.train_loader):]) / len(self.grad_norms),184 "total_params": sum(p.numel() for p in self.model.parameters()),185 "latency_ms": latency_est,186 "energy_uj": energy_est,187 "time_s": elapsed,188 }189 190 # Extract TT stats191 if hasattr(self.model, "stats"):192 metrics["model_stats"] = self.model.stats193 194 # Validation195 if self.val_loader is not None:196 val_metrics = self.validate()197 metrics.update(val_metrics)198 199 self.metrics_history.append(metrics)200 return metrics201 202 @torch.no_grad()203 def validate(self) -> Dict:204 """Run validation."""205 self.model.eval()206 total_loss = 0.0207 total_tokens = 0208 209 for inputs, targets in self.val_loader:210 inputs, targets = inputs.to(self.device), targets.to(self.device)211 logits = self.model(inputs)212 loss = F.cross_entropy(213 logits.reshape(-1, logits.size(-1)),214 targets.reshape(-1),215 ignore_index=0,216 reduction="sum",217 )218 total_loss += loss.item()219 total_tokens += inputs.numel()220 221 avg_loss = total_loss / max(total_tokens, 1)222 return {223 "val_loss": avg_loss,224 "val_ppl": math.exp(min(avg_loss, 20.0)),225 }226 227 @torch.no_grad()228 def evaluate(self) -> Dict:229 """230 Full evaluation on test set.231 Returns comprehensive metrics dict.232 """233 self.model.eval()234 total_loss = 0.0235 total_tokens = 0236 latency_samples = []237 238 for inputs, targets in self.test_loader:239 inputs, targets = inputs.to(self.device), targets.to(self.device)240 241 t0 = time.time()242 logits = self.model(inputs)243 t1 = time.time()244 latency_samples.append((t1 - t0) * 1000 / inputs.size(0)) # ms per sample245 246 loss = F.cross_entropy(247 logits.reshape(-1, logits.size(-1)),248 targets.reshape(-1),249 ignore_index=0,250 reduction="sum",251 )252 total_loss += loss.item()253 total_tokens += inputs.numel()254 255 avg_loss = total_loss / max(total_tokens, 1)256 257 return {258 "test_loss": avg_loss,259 "test_ppl": math.exp(min(avg_loss, 20.0)),260 "latency_ms_mean": sum(latency_samples) / len(latency_samples),261 "total_params": self.model.total_params,262 "energy_uj": self.energy_estimator.estimate(self.model),263 "model_stats": getattr(self.model, "stats", {}),264 }265 266 def train(self) -> Dict:267 """Full training loop."""268 best_val_ppl = float("inf")269 270 for epoch in range(self.config.training.max_epochs):271 metrics = self.train_epoch(epoch)272 273 # Logging274 print(f"Epoch {epoch+1}/{self.config.training.max_epochs}: "275 f"train_ppl={metrics['train_ppl']:.2f} "276 f"val_ppl={metrics.get('val_ppl', 'N/A')} "277 f"lr={metrics['lr']:.2e}")278 279 if metrics.get("val_ppl", float("inf")) < best_val_ppl:280 best_val_ppl = metrics["val_ppl"]281 self.save_checkpoint("best")282 283 # Early stopping checks284 if self.budget_tracker.exceeds_budget(metrics, self.config.model):285 print(f"[BUDGET] Exceeded constraints. Stopping.")286 break287 288 self.save_checkpoint("last")289 self.save_metrics()290 return self.metrics_history[-1] if self.metrics_history else {}291 292 def save_checkpoint(self, tag: str = "checkpoint"):293 """Save model checkpoint with metadata."""294 self.output_dir.mkdir(parents=True, exist_ok=True)295 path = self.output_dir / f"{tag}.pt"296 torch.save({297 "model_state_dict": self.model.state_dict(),298 "optimizer_state_dict": self.optimizer.state_dict(),299 "config": self.config,300 "metrics": self.metrics_history,301 }, path)302 print(f"Checkpoint saved to {path}")303 304 def load_checkpoint(self, tag: str = "best"):305 """Load checkpoint."""306 path = self.output_dir / f"{tag}.pt"307 if not path.exists():308 print(f"Checkpoint {path} not found")309 return310 ckpt = torch.load(path, map_location=self.device, weights_only=True)311 self.model.load_state_dict(ckpt["model_state_dict"])312 self.optimizer.load_state_dict(ckpt["optimizer_state_dict"])313 314 def save_metrics(self):315 """Save metrics to JSON."""316 self.output_dir.mkdir(parents=True, exist_ok=True)317 path = self.output_dir / "metrics.json"318 with open(path, "w") as f:319 json.dump(self.metrics_history, f, indent=2)320 print(f"Metrics saved to {path}")321 322 323class DistillationTrainer(Trainer):324 """325 Knowledge distillation trainer.326 327 Student = compressed Q-TensorFormer.328 Teacher = dense (or larger) model.329 """330 331 def __init__(self, student: nn.Module, teacher: nn.Module, *args,332 alpha: float = 0.5, temperature: float = 3.0, **kwargs):333 """334 Args:335 student: Compressed Q-TensorFormer.336 teacher: Dense baseline (frozen).337 alpha: Weight between distillation loss (α) and task loss (1-α).338 temperature: Softmax temperature.339 """340 super().__init__(student, *args, **kwargs)341 self.teacher = teacher.to(self.device)342 self.teacher.eval()343 self.alpha = alpha344 self.temperature = temperature345 346 # Freeze teacher347 for p in self.teacher.parameters():348 p.requires_grad = False349 350 def train_epoch(self, epoch: int) -> Dict:351 self.model.train()352 total_loss = 0.0353 total_tokens = 0354 355 for step, (inputs, targets) in enumerate(self.train_loader):356 inputs, targets = inputs.to(self.device), targets.to(self.device)357 358 self.optimizer.zero_grad()359 360 # Student forward361 logits, stats = self.model(inputs, return_stats=True)362 363 # Task loss364 task_loss = F.cross_entropy(365 logits.reshape(-1, logits.size(-1)),366 targets.reshape(-1),367 ignore_index=0,368 )369 370 # Distillation loss371 with torch.no_grad():372 teacher_logits = self.teacher(inputs)373 374 distill_loss = F.kl_div(375 F.log_softmax(logits / self.temperature, dim=-1),376 F.softmax(teacher_logits / self.temperature, dim=-1),377 reduction="batchmean",378 ) * (self.temperature ** 2)379 380 loss = (1 - self.alpha) * task_loss + self.alpha * distill_loss381 loss.backward()382 383 torch.nn.utils.clip_grad_norm_(384 self.model.parameters(), self.config.training.max_grad_norm385 )386 self.optimizer.step()387 self.scheduler.step()388 389 total_loss += task_loss.item() * inputs.numel()390 total_tokens += inputs.numel()391 392 avg_loss = total_loss / max(total_tokens, 1)393 ppl = math.exp(min(avg_loss, 20.0))394 return {395 "epoch": epoch,396 "train_loss": avg_loss,397 "train_ppl": ppl,398 "lr": self.optimizer.param_groups[0]["lr"],399 }400 