CoolFace
Modelpublic

Premchan369/Q-TensorFormer

sourceHugging Faceapache-2.0updated 9d agoView on Hugging Face
2likes185downloads
training.py400 linesDownload Raw Back to src
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