amachine.am_transformers.am_trainer
1import json 2import math 3import time 4 5from pathlib import Path 6from typing import Tuple, Any 7from datetime import datetime, timezone 8 9import warnings 10import torch 11 12torch.compiler.config.force_disable_caches = True 13 14from torch.utils.data import DataLoader 15from packaging.version import Version 16 17import pytorch_optimizer as pt_optim 18 19from tokenizers import Tokenizer 20 21from transformers import ( 22 AutoConfig, 23 AutoModelForCausalLM 24) 25 26from .am_training_config import TrainingConfig, TrainerState, TrainerContext 27from .am_trainer_utils import get_lr 28from .am_datasets import StreamingParquetDataset, build_remap_table 29from .am_text_datasets import TextStreamingParquetDataset 30 31from .am_metrics import GradNormTracker, Accumulator 32from .am_eval import run_eval 33from .am_logger import log_metrics 34from .am_checkpointer import CheckpointManager 35from .am_set_trainable import set_trainable_parameters 36 37try: 38 from .am_control_model_exp import * 39except Exception: 40 import traceback 41 traceback.print_exc() 42 warnings.warn("Failed to import control model") 43 44class Trainer: 45 46 def __init__(self, config: TrainingConfig): 47 48 from aim import Run 49 50 self.training_config = config 51 52 # Environment abd Seeding 53 torch.manual_seed(self.training_config.seed) 54 if torch.cuda.is_available(): 55 torch.cuda.manual_seed_all(self.training_config.seed) 56 torch.set_float32_matmul_precision('high') 57 58 Path(self.training_config.output_dir).mkdir(exist_ok=True) 59 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") 60 61 # Save config early so we have a record even if initialization crashes 62 tr_config_path = Path(self.training_config.output_dir) / "training_config.json" 63 self.training_config.save(tr_config_path) 64 65 self.fused_adam = ( 66 self.device.type == "cuda" 67 and Version(torch.__version__.split("+")[0]) >= Version("2.0") 68 ) 69 self.ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.pdtype] 70 self.wdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.wdtype] 71 72 self.use_amp = self.ptdtype != torch.float32 73 self.min_lr = self.training_config.lr * self.training_config.min_lr_ratio 74 75 # Create default state counters (epoch=0, step=0, etc.) 76 self.trainer_state = TrainerState() 77 self.aim_run_hash = None 78 79 # Build the model, optimizer, and DataLoaders unconditionally. 80 self._build_components() 81 self._save_runtime_info(Path(self.training_config.output_dir)) 82 83 self.checkpointer = CheckpointManager( 84 config=self.training_config, 85 model=self.model, 86 optimizer=self.optimizer, 87 trainer_state=self.trainer_state, 88 device=self.device 89 ) 90 91 if self.training_config.resume is not None: 92 93 # This updates model, optimizer, and trainer_state in place. 94 self.checkpointer.load(self.training_config.resume) 95 self.aim_run_hash = self.trainer_state.aim_run_hash 96 97 # Fast-forward the dataset iterator if we aren't starting fresh 98 if not self.training_config.resume_reset_data: 99 self._fast_forward_dataloader(self.trainer_state.batches_consumed) 100 101 # Initialized after the checkpoint loads 102 self.aim_run = None 103 if self.training_config.aim_repo is not None: 104 105 experiment = self.training_config.experiment if self.training_config.experiment else "Train" 106 107 self.aim_run = Run( 108 repo=self.training_config.aim_repo, 109 experiment=experiment, 110 log_system_params=True, 111 run_hash=self.aim_run_hash, 112 ) 113 114 self.aim_run["hparams"] = self.training_config.to_dict() 115 if self.training_config.experiment_params is not None: 116 self.aim_run["experiment_params"] = self.training_config.experiment_params 117 118 self.ctx = TrainerContext( 119 model=self.model, 120 optimizer=self.optimizer, 121 device=self.device, 122 ptdtype=self.ptdtype, 123 use_amp=self.use_amp, 124 pad_id=self.pad_id, 125 eos_id=self.eos_id, 126 training_config=self.training_config, 127 aim_run=self.aim_run 128 ) 129 130 # Helpers 131 132 def _runtime_info(self) -> dict[str, Any]: 133 info = { 134 "created_at": datetime.now(timezone.utc).isoformat(), 135 "pytorch_version": torch.__version__, 136 "device": str(self.device) 137 } 138 return info 139 140 def _save_runtime_info(self, dirpath: Path) -> None: 141 with open(dirpath / "runtime_info.json", "w") as f: 142 json.dump(self._runtime_info(), f, indent=2) 143 144 def _unwrap_model(self): 145 """Return the original un-compiled model for serialisation.""" 146 return self.model._orig_mod if hasattr(self.model, "_orig_mod") else self.model 147 148 def _configure_model(self, model): 149 150 """Apply gradient checkpointing and torch.compile. Returns the configured model.""" 151 model.config.use_cache = False 152 153 if not self.training_config.no_grad_ckpt: 154 model.gradient_checkpointing_enable( 155 gradient_checkpointing_kwargs={"use_reentrant": False} 156 ) 157 print("Gradient checkpointing: ON (use_reentrant=False)") 158 159 if self.training_config.compile: 160 print("Compiling model with torch.compile … (first step will be slow)") 161 model = torch.compile(model) 162 163 return model 164 165 def _fast_forward_dataloader(self, n_batches: int) -> None: 166 167 if n_batches <= 0: 168 return 169 170 print(f" Fast-forwarding {n_batches:,} micro-batches …") 171 t_ff = time.time() 172 173 # Replay from epoch 0, letting epoch transitions happen naturally 174 replay_epoch = 0 175 self.loader.dataset.set_epoch(replay_epoch) 176 self.data_iter = iter(self.loader) 177 178 for i in range(n_batches): 179 try: 180 next(self.data_iter) 181 except StopIteration: 182 replay_epoch += 1 183 self.loader.dataset.set_epoch(replay_epoch) 184 self.data_iter = iter(self.loader) 185 next(self.data_iter) 186 187 if (i + 1) % 10_000 == 0: 188 elapsed_ff = time.time() - t_ff 189 rate = (i + 1) / max(elapsed_ff, 1e-6) 190 eta = (n_batches - i - 1) / max(rate, 1e-6) 191 print(f" … {i+1:,} / {n_batches:,} ({rate:.0f} batches/s, ETA {eta:.0f}s)") 192 193 print(f" Fast-forward complete in {time.time() - t_ff:.1f}s") 194 195 def _get_next_batch(self): 196 197 try: 198 batch = next(self.data_iter) 199 except StopIteration: 200 self.trainer_state.epoch += 1 201 self.loader.dataset.set_epoch( self.trainer_state.epoch ) 202 print(f" [Data] Epoch {self.trainer_state.epoch}, restarting iterator.") 203 self.data_iter = iter(self.loader) 204 batch = next(self.data_iter) 205 206 self.trainer_state.batches_consumed += 1 # track total micro-batches for resume 207 return batch 208 209 # Component construction 210 211 def _build_components(self): 212 213 pretokenized : bool = self.training_config.pretokenized 214 215 model_dir = Path(self.training_config.model_dir) 216 self.tokenizer = Tokenizer.from_file(str(model_dir / "tokenizer.json")) 217 218 with open(model_dir / "config.json", "r") as f: 219 model_config = json.load(f) 220 221 ######################################################################### 222 223 def extract_token_str(token_val): 224 """Helper to handle Hugging Face tokenizer config dicts.""" 225 if isinstance(token_val, dict): 226 return token_val.get("content", None) 227 return token_val 228 229 self.unk_token = None 230 self.pad_token = None 231 self.eos_token = None 232 233 if (model_dir / "tokenizer_config.json").exists(): 234 with open(model_dir / "tokenizer_config.json", "r") as f: 235 tokenizer_config = json.load(f) 236 237 # Safely extract string values, even if they are dictionaries 238 self.unk_token = extract_token_str(tokenizer_config.get('unk_token')) 239 self.pad_token = extract_token_str(tokenizer_config.get('pad_token')) 240 self.eos_token = extract_token_str(tokenizer_config.get('eos_token')) 241 242 # Get unk_id (accounting for None if not found) 243 self.unk_id = self.tokenizer.token_to_id(self.unk_token) if self.unk_token else None 244 self.pad_id = model_config.get("pad_token_id") 245 self.eos_id = model_config.get("eos_token_id") 246 247 assert self.pad_id is not None, "pad_token_id is missing from config.json" 248 assert self.eos_id is not None, "eos_token_id is missing from config.json" 249 250 # Validate PAD 251 if self.pad_token is not None: 252 actual_pad_id = self.tokenizer.token_to_id(self.pad_token) 253 if actual_pad_id is None: 254 raise ValueError(f"pad_token '{self.pad_token}' not found in tokenizer vocabulary.") 255 if actual_pad_id != self.pad_id: 256 raise ValueError(f"Mismatch: Tokenizer pad_id ({actual_pad_id}) != Model pad_id ({self.pad_id})") 257 258 # Validate EOS 259 if self.eos_token is not None: 260 actual_eos_id = self.tokenizer.token_to_id(self.eos_token) 261 if actual_eos_id is None: 262 raise ValueError(f"eos_token '{self.eos_token}' not found in tokenizer vocabulary.") 263 if actual_eos_id != self.eos_id: 264 raise ValueError(f"Mismatch: Tokenizer eos_id ({actual_eos_id}) != Model eos_id ({self.eos_id})") 265 266 ########################################################################## 267 268 def make_loader(path: str) -> DataLoader: 269 270 if self.training_config.tokenizer_type == "symbolic": 271 272 assert isinstance( self.training_config.metadata, str ) 273 remap = build_remap_table( 274 self.training_config.metadata, 275 self.tokenizer, 276 unk_token=self.unk_token 277 ) 278 279 ds = StreamingParquetDataset( 280 path=path, 281 input_column=self.training_config.input_col, 282 remap=remap, 283 seq_len=self.training_config.seq_len, 284 shuffle_buffer_size=self.training_config.shuffle_buffer 285 ) 286 287 elif self.training_config.tokenizer_type == "text": 288 289 ds = TextStreamingParquetDataset( 290 path=path, 291 input_column=self.training_config.input_col, 292 tokenizer_path=str(model_dir / "tokenizer.json"), 293 seq_len=self.training_config.seq_len, 294 eos_id=self.eos_id, 295 end_docs_with_eos=( None if pretokenized else True ), 296 pretokenized=pretokenized, 297 shuffle_buffer_size=self.training_config.shuffle_buffer 298 ) 299 else: 300 raise ValueError(f"Unknown tokenizer_type: {self.training_config.tokenizer_type}") 301 302 g = torch.Generator() 303 g.manual_seed( self.training_config.seed ) 304 305 return DataLoader( 306 ds, 307 batch_size=self.training_config.batch_size, 308 num_workers=self.training_config.num_workers, 309 pin_memory=(self.device.type == "cuda"), 310 persistent_workers=(self.training_config.num_workers > 0), 311 shuffle=False, 312 drop_last=True, 313 generator=g, 314 in_order=True, 315 prefetch_factor=3, 316 ) 317 318 self.loader = make_loader(self.training_config.data) 319 self.data_iter = iter(self.loader) 320 321 self.eval_loader = make_loader( 322 self.training_config.eval_data 323 ) if self.training_config.eval_data else None 324 325 model_config = AutoConfig.from_pretrained(self.training_config.model_dir) 326 327 kwargs = {} 328 if self.training_config.attn_implementation is not None: 329 kwargs["attn_implementation"] = self.training_config.attn_implementation 330 try: 331 self.model = AutoModelForCausalLM.from_config( 332 model_config, 333 **kwargs 334 ).to(device=self.device, dtype=self.wdtype) 335 336 except ValueError as e: 337 raise e 338 339 self.model = self._configure_model(self.model) 340 341 trainable_parameters = set_trainable_parameters( 342 model=self.model, 343 layers_to_train=self.training_config.train_layers, 344 train_input_embeddings=self.training_config.train_input_embeddings, 345 train_output_layer=self.training_config.train_output_layer, 346 strict_tie_check=True, 347 verbose=True, 348 ) 349 350 if self.training_config.optimizer == "adamw": 351 self.optimizer = torch.optim.AdamW( 352 trainable_parameters, 353 lr=self.training_config.lr, 354 betas=(self.training_config.beta1, self.training_config.beta2), 355 eps=self.training_config.eps, 356 weight_decay=self.training_config.weight_decay, 357 fused=self.fused_adam, 358 ) 359 elif self.training_config.optimizer == "sgd": 360 self.optimizer = torch.optim.SGD( 361 trainable_parameters, 362 lr=self.training_config.lr, 363 momentum=self.training_config.momentum, 364 weight_decay=self.training_config.weight_decay, 365 nesterov=True, 366 fused=self.fused_adam, 367 ) 368 elif self.training_config.optimizer == "novograd": 369 self.optimizer = pt_optim.NovoGrad( 370 trainable_parameters, 371 lr=self.training_config.lr, 372 betas=(self.training_config.beta1, self.training_config.beta2), 373 eps=self.training_config.eps, 374 weight_decay=self.training_config.weight_decay, 375 weight_decouple=True 376 ) 377 elif self.training_config.optimizer == "fromage": 378 self.optimizer = pt_optim.Fromage( 379 trainable_parameters, 380 lr=self.training_config.lr, 381 p_bound=None, 382 ) 383 elif self.training_config.optimizer == "lamb": 384 self.optimizer = pt_optim.Lamb( 385 trainable_parameters, 386 lr=self.training_config.lr, 387 betas=(self.training_config.beta1, self.training_config.beta2), 388 eps=self.training_config.eps, 389 weight_decay=self.training_config.weight_decay, 390 weight_decouple=True, 391 rectify=True, 392 max_grad_norm=1.0, 393 ) 394 elif self.training_config.optimizer == "nero": 395 self.optimizer = pt_optim.Nero( 396 trainable_parameters, 397 lr=self.training_config.lr, 398 eps=self.training_config.eps, 399 ) 400 elif self.training_config.optimizer == "scion": 401 self.optimizer = pt_optim.SCION( 402 trainable_parameters, 403 lr=self.training_config.lr, 404 weight_decay=self.training_config.weight_decay 405 ) 406 elif self.training_config.optimizer == "apollo": 407 self.optimizer = pt_optim.APOLLO( 408 trainable_parameters, 409 lr=self.training_config.lr, 410 weight_decay=self.training_config.weight_decay 411 ) 412 else: 413 raise ValueError(f"Unsupported optimizer {self.training_config.optimizer}") 414 415 n_params = sum(p.numel() for p in self._unwrap_model().parameters()) 416 print(f"Model parameters: {n_params:,} | device: {self.device} | dtype: {self.training_config.pdtype}") 417 418 self.grad_norm_tracker = GradNormTracker(self._unwrap_model(), self.device) 419 420 # Save Checkpoint 421 422 def _checkpoint_if_needed(self) -> bool: 423 if self.trainer_state.global_step % self.training_config.save_every == 0: 424 self.checkpointer.save(self.trainer_state.global_step) 425 return True 426 return False 427 428 def save_checkpoint(self, step: int): 429 self.checkpointer.save(step) 430 431 # Evaluate 432 433 def _evaluate_if_needed(self) -> dict | None: 434 435 if self.eval_loader is None: 436 return None 437 438 if self.trainer_state.global_step % self.training_config.eval_every != 0: 439 return None 440 441 return run_eval( 442 eval_loader=self.eval_loader, 443 model=self.ctx.model, 444 training_config=self.ctx.training_config, 445 device=self.ctx.device, 446 ptdtype=self.ctx.ptdtype, 447 use_amp=self.ctx.use_amp, 448 pad_id=self.pad_id, 449 pad_token=self.pad_token, 450 eos_token=self.eos_token, 451 unk_token=self.unk_token, 452 tokenizer=self.tokenizer, 453 ) 454 455 # Set Learning Rate 456 457 def _set_lr(self, step: int) -> float: 458 lr = get_lr(step, self.training_config, self.min_lr) 459 for pg in self.optimizer.param_groups: 460 pg["lr"] = lr 461 return lr 462 463 # Training step 464 465 def training_step(self) -> Tuple[float, int, float, bool, tuple[dict,dict], dict ]: 466 467 raw_loss_acc = torch.zeros(1, device=self.device) 468 token_acc = torch.zeros(1, device=self.device) 469 470 extra_term_names = self.training_config.extra_loss_terms or [] 471 extra_term_acc = {name: torch.zeros(1, device=self.device) for name in extra_term_names} 472 473 self.optimizer.zero_grad(set_to_none=True) 474 475 # Forward and Backward Passes 476 477 for _ in range(self.training_config.grad_accum): 478 479 batch = self._get_next_batch() 480 481 input_ids = batch["input_ids"].to(self.device, non_blocking=True) 482 attention_mask = batch["attention_mask"].to(self.device, non_blocking=True) 483 484 labels = input_ids.masked_fill(attention_mask == 0, -100) 485 n_tokens = attention_mask[..., 1:].sum() 486 487 token_acc += n_tokens 488 489 with torch.autocast(device_type=self.device.type, dtype=self.ptdtype, enabled=self.use_amp): 490 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) 491 loss_scaled = outputs.loss * n_tokens 492 493 loss_scaled.backward() 494 raw_loss_acc += loss_scaled.detach() 495 496 for name in extra_term_names: 497 term_val = getattr( outputs, name, None ) 498 if term_val is not None: 499 extra_term_acc[name] += term_val 500 501 # Single CPU-GPU sync for both accumulators 502 raw_loss_sum_weighted, accum_tokens = raw_loss_acc.item(), token_acc.item() 503 step_tokens = int(accum_tokens) 504 extra_terms_sum = {name: acc.item() / self.training_config.grad_accum for name, acc in extra_term_acc.items()} 505 506 # Gradient Normalization and Stability Check 507 if accum_tokens > 0: 508 grads = [p.grad for p in self.model.parameters() if p.grad is not None] 509 torch._foreach_div_(grads, accum_tokens) 510 511 if not math.isfinite(raw_loss_sum_weighted): 512 print(f"Numerical instability at Step {self.trainer_state.global_step}. Exiting") 513 raise RuntimeError(f"Loss is {raw_loss_sum_weighted} (NaN/Inf).") 514 515 # Spike Detection and EMA Updates 516 token_weighted_step_mean_loss = raw_loss_sum_weighted / max(accum_tokens, 1) 517 is_warm = self.trainer_state.global_step >= self.training_config.spike_warmup_steps 518 519 is_spike = self.trainer_state.process_loss( 520 token_weighted_step_mean_loss, 521 self.training_config, 522 is_warm) 523 524 if is_spike and False: 525 526 self.optimizer.zero_grad(set_to_none=True) 527 if self.trainer_state.consecutive_spikes >= self.trainer_state.max_consecutive_spikes: 528 529 # self.save_checkpoint(self.trainer_state.global_step) 530 # raise RuntimeError( 531 # f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 532 # f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 533 # ) 534 535 print( 536 f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 537 f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 538 ) 539 540 print( f"[spike] loss={token_weighted_step_mean_loss:.4f} vs ema={self.trainer_state.spike_loss_ema:.4f}" ) 541 return raw_loss_sum_weighted, step_tokens, -1.0, True, ({}, {}), {} 542 543 if self.training_config.track_grad_norms: 544 agg_norms, per_layer_norms = self.grad_norm_tracker.compute() 545 else: 546 agg_norms, per_layer_norms = {}, {} 547 548 max_norm = self.training_config.grad_clip if self.training_config.grad_clip > 0.0 else float('inf') 549 grad_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm).item() 550 551 self.optimizer.step() 552 553 return raw_loss_sum_weighted, step_tokens, grad_norm, False, (agg_norms, per_layer_norms), extra_terms_sum 554 555 # Main Training loop 556 557 def train(self): 558 559 if self.trainer_state.global_step >= self.training_config.steps: 560 raise ValueError(f"global_step ({self.trainer_state.global_step}) >= steps ({self.training_config.steps}).") 561 562 self.model.train() 563 accumulator = Accumulator() 564 t0 = time.time() 565 566 for step in range(self.trainer_state.global_step, self.training_config.steps): 567 568 # Step 569 lr_now = self._set_lr(step) 570 step_outputs = self.training_step() 571 572 # Accumulate State 573 accumulator.update(*step_outputs) 574 575 self.trainer_state.global_step += 1 576 self.trainer_state.global_tokens += step_outputs[1] 577 578 # Log and Evaluate 579 if self.trainer_state.global_step % self.training_config.log_every == 0: 580 581 metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 582 eval_metrics = self._evaluate_if_needed() 583 584 log_metrics(self.ctx, self.trainer_state, metrics, eval_metrics) 585 586 accumulator = Accumulator() 587 t0 = time.time() 588 589 # Checkpoint 590 if self._checkpoint_if_needed(): 591 t0 = time.time() 592 593 # Done Training, Finalize 594 if accumulator.windowed_steps > 0: 595 596 final_metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 597 598 # Force a final eval if we have a loader 599 final_eval = run_eval( 600 eval_loader=self.eval_loader, 601 model=self.ctx.model, 602 training_config=self.ctx.training_config, 603 device=self.ctx.device, 604 ptdtype=self.ctx.ptdtype, 605 use_amp=self.ctx.use_amp, 606 pad_id=self.pad_id, 607 pad_token=self.pad_token, 608 eos_token=self.eos_token, 609 unk_token=self.unk_token, 610 tokenizer=self.tokenizer, 611 ) if self.eval_loader is not None else None 612 613 log_metrics(self.ctx, self.trainer_state, final_metrics, final_eval) 614 615 if self.aim_run is not None: 616 self.aim_run.close() 617 618 self.save_checkpoint(self.trainer_state.global_step) 619 print(f"Training complete. Final model: step_{self.trainer_state.global_step:07d}")
class
Trainer:
45class Trainer: 46 47 def __init__(self, config: TrainingConfig): 48 49 from aim import Run 50 51 self.training_config = config 52 53 # Environment abd Seeding 54 torch.manual_seed(self.training_config.seed) 55 if torch.cuda.is_available(): 56 torch.cuda.manual_seed_all(self.training_config.seed) 57 torch.set_float32_matmul_precision('high') 58 59 Path(self.training_config.output_dir).mkdir(exist_ok=True) 60 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") 61 62 # Save config early so we have a record even if initialization crashes 63 tr_config_path = Path(self.training_config.output_dir) / "training_config.json" 64 self.training_config.save(tr_config_path) 65 66 self.fused_adam = ( 67 self.device.type == "cuda" 68 and Version(torch.__version__.split("+")[0]) >= Version("2.0") 69 ) 70 self.ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.pdtype] 71 self.wdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.wdtype] 72 73 self.use_amp = self.ptdtype != torch.float32 74 self.min_lr = self.training_config.lr * self.training_config.min_lr_ratio 75 76 # Create default state counters (epoch=0, step=0, etc.) 77 self.trainer_state = TrainerState() 78 self.aim_run_hash = None 79 80 # Build the model, optimizer, and DataLoaders unconditionally. 81 self._build_components() 82 self._save_runtime_info(Path(self.training_config.output_dir)) 83 84 self.checkpointer = CheckpointManager( 85 config=self.training_config, 86 model=self.model, 87 optimizer=self.optimizer, 88 trainer_state=self.trainer_state, 89 device=self.device 90 ) 91 92 if self.training_config.resume is not None: 93 94 # This updates model, optimizer, and trainer_state in place. 95 self.checkpointer.load(self.training_config.resume) 96 self.aim_run_hash = self.trainer_state.aim_run_hash 97 98 # Fast-forward the dataset iterator if we aren't starting fresh 99 if not self.training_config.resume_reset_data: 100 self._fast_forward_dataloader(self.trainer_state.batches_consumed) 101 102 # Initialized after the checkpoint loads 103 self.aim_run = None 104 if self.training_config.aim_repo is not None: 105 106 experiment = self.training_config.experiment if self.training_config.experiment else "Train" 107 108 self.aim_run = Run( 109 repo=self.training_config.aim_repo, 110 experiment=experiment, 111 log_system_params=True, 112 run_hash=self.aim_run_hash, 113 ) 114 115 self.aim_run["hparams"] = self.training_config.to_dict() 116 if self.training_config.experiment_params is not None: 117 self.aim_run["experiment_params"] = self.training_config.experiment_params 118 119 self.ctx = TrainerContext( 120 model=self.model, 121 optimizer=self.optimizer, 122 device=self.device, 123 ptdtype=self.ptdtype, 124 use_amp=self.use_amp, 125 pad_id=self.pad_id, 126 eos_id=self.eos_id, 127 training_config=self.training_config, 128 aim_run=self.aim_run 129 ) 130 131 # Helpers 132 133 def _runtime_info(self) -> dict[str, Any]: 134 info = { 135 "created_at": datetime.now(timezone.utc).isoformat(), 136 "pytorch_version": torch.__version__, 137 "device": str(self.device) 138 } 139 return info 140 141 def _save_runtime_info(self, dirpath: Path) -> None: 142 with open(dirpath / "runtime_info.json", "w") as f: 143 json.dump(self._runtime_info(), f, indent=2) 144 145 def _unwrap_model(self): 146 """Return the original un-compiled model for serialisation.""" 147 return self.model._orig_mod if hasattr(self.model, "_orig_mod") else self.model 148 149 def _configure_model(self, model): 150 151 """Apply gradient checkpointing and torch.compile. Returns the configured model.""" 152 model.config.use_cache = False 153 154 if not self.training_config.no_grad_ckpt: 155 model.gradient_checkpointing_enable( 156 gradient_checkpointing_kwargs={"use_reentrant": False} 157 ) 158 print("Gradient checkpointing: ON (use_reentrant=False)") 159 160 if self.training_config.compile: 161 print("Compiling model with torch.compile … (first step will be slow)") 162 model = torch.compile(model) 163 164 return model 165 166 def _fast_forward_dataloader(self, n_batches: int) -> None: 167 168 if n_batches <= 0: 169 return 170 171 print(f" Fast-forwarding {n_batches:,} micro-batches …") 172 t_ff = time.time() 173 174 # Replay from epoch 0, letting epoch transitions happen naturally 175 replay_epoch = 0 176 self.loader.dataset.set_epoch(replay_epoch) 177 self.data_iter = iter(self.loader) 178 179 for i in range(n_batches): 180 try: 181 next(self.data_iter) 182 except StopIteration: 183 replay_epoch += 1 184 self.loader.dataset.set_epoch(replay_epoch) 185 self.data_iter = iter(self.loader) 186 next(self.data_iter) 187 188 if (i + 1) % 10_000 == 0: 189 elapsed_ff = time.time() - t_ff 190 rate = (i + 1) / max(elapsed_ff, 1e-6) 191 eta = (n_batches - i - 1) / max(rate, 1e-6) 192 print(f" … {i+1:,} / {n_batches:,} ({rate:.0f} batches/s, ETA {eta:.0f}s)") 193 194 print(f" Fast-forward complete in {time.time() - t_ff:.1f}s") 195 196 def _get_next_batch(self): 197 198 try: 199 batch = next(self.data_iter) 200 except StopIteration: 201 self.trainer_state.epoch += 1 202 self.loader.dataset.set_epoch( self.trainer_state.epoch ) 203 print(f" [Data] Epoch {self.trainer_state.epoch}, restarting iterator.") 204 self.data_iter = iter(self.loader) 205 batch = next(self.data_iter) 206 207 self.trainer_state.batches_consumed += 1 # track total micro-batches for resume 208 return batch 209 210 # Component construction 211 212 def _build_components(self): 213 214 pretokenized : bool = self.training_config.pretokenized 215 216 model_dir = Path(self.training_config.model_dir) 217 self.tokenizer = Tokenizer.from_file(str(model_dir / "tokenizer.json")) 218 219 with open(model_dir / "config.json", "r") as f: 220 model_config = json.load(f) 221 222 ######################################################################### 223 224 def extract_token_str(token_val): 225 """Helper to handle Hugging Face tokenizer config dicts.""" 226 if isinstance(token_val, dict): 227 return token_val.get("content", None) 228 return token_val 229 230 self.unk_token = None 231 self.pad_token = None 232 self.eos_token = None 233 234 if (model_dir / "tokenizer_config.json").exists(): 235 with open(model_dir / "tokenizer_config.json", "r") as f: 236 tokenizer_config = json.load(f) 237 238 # Safely extract string values, even if they are dictionaries 239 self.unk_token = extract_token_str(tokenizer_config.get('unk_token')) 240 self.pad_token = extract_token_str(tokenizer_config.get('pad_token')) 241 self.eos_token = extract_token_str(tokenizer_config.get('eos_token')) 242 243 # Get unk_id (accounting for None if not found) 244 self.unk_id = self.tokenizer.token_to_id(self.unk_token) if self.unk_token else None 245 self.pad_id = model_config.get("pad_token_id") 246 self.eos_id = model_config.get("eos_token_id") 247 248 assert self.pad_id is not None, "pad_token_id is missing from config.json" 249 assert self.eos_id is not None, "eos_token_id is missing from config.json" 250 251 # Validate PAD 252 if self.pad_token is not None: 253 actual_pad_id = self.tokenizer.token_to_id(self.pad_token) 254 if actual_pad_id is None: 255 raise ValueError(f"pad_token '{self.pad_token}' not found in tokenizer vocabulary.") 256 if actual_pad_id != self.pad_id: 257 raise ValueError(f"Mismatch: Tokenizer pad_id ({actual_pad_id}) != Model pad_id ({self.pad_id})") 258 259 # Validate EOS 260 if self.eos_token is not None: 261 actual_eos_id = self.tokenizer.token_to_id(self.eos_token) 262 if actual_eos_id is None: 263 raise ValueError(f"eos_token '{self.eos_token}' not found in tokenizer vocabulary.") 264 if actual_eos_id != self.eos_id: 265 raise ValueError(f"Mismatch: Tokenizer eos_id ({actual_eos_id}) != Model eos_id ({self.eos_id})") 266 267 ########################################################################## 268 269 def make_loader(path: str) -> DataLoader: 270 271 if self.training_config.tokenizer_type == "symbolic": 272 273 assert isinstance( self.training_config.metadata, str ) 274 remap = build_remap_table( 275 self.training_config.metadata, 276 self.tokenizer, 277 unk_token=self.unk_token 278 ) 279 280 ds = StreamingParquetDataset( 281 path=path, 282 input_column=self.training_config.input_col, 283 remap=remap, 284 seq_len=self.training_config.seq_len, 285 shuffle_buffer_size=self.training_config.shuffle_buffer 286 ) 287 288 elif self.training_config.tokenizer_type == "text": 289 290 ds = TextStreamingParquetDataset( 291 path=path, 292 input_column=self.training_config.input_col, 293 tokenizer_path=str(model_dir / "tokenizer.json"), 294 seq_len=self.training_config.seq_len, 295 eos_id=self.eos_id, 296 end_docs_with_eos=( None if pretokenized else True ), 297 pretokenized=pretokenized, 298 shuffle_buffer_size=self.training_config.shuffle_buffer 299 ) 300 else: 301 raise ValueError(f"Unknown tokenizer_type: {self.training_config.tokenizer_type}") 302 303 g = torch.Generator() 304 g.manual_seed( self.training_config.seed ) 305 306 return DataLoader( 307 ds, 308 batch_size=self.training_config.batch_size, 309 num_workers=self.training_config.num_workers, 310 pin_memory=(self.device.type == "cuda"), 311 persistent_workers=(self.training_config.num_workers > 0), 312 shuffle=False, 313 drop_last=True, 314 generator=g, 315 in_order=True, 316 prefetch_factor=3, 317 ) 318 319 self.loader = make_loader(self.training_config.data) 320 self.data_iter = iter(self.loader) 321 322 self.eval_loader = make_loader( 323 self.training_config.eval_data 324 ) if self.training_config.eval_data else None 325 326 model_config = AutoConfig.from_pretrained(self.training_config.model_dir) 327 328 kwargs = {} 329 if self.training_config.attn_implementation is not None: 330 kwargs["attn_implementation"] = self.training_config.attn_implementation 331 try: 332 self.model = AutoModelForCausalLM.from_config( 333 model_config, 334 **kwargs 335 ).to(device=self.device, dtype=self.wdtype) 336 337 except ValueError as e: 338 raise e 339 340 self.model = self._configure_model(self.model) 341 342 trainable_parameters = set_trainable_parameters( 343 model=self.model, 344 layers_to_train=self.training_config.train_layers, 345 train_input_embeddings=self.training_config.train_input_embeddings, 346 train_output_layer=self.training_config.train_output_layer, 347 strict_tie_check=True, 348 verbose=True, 349 ) 350 351 if self.training_config.optimizer == "adamw": 352 self.optimizer = torch.optim.AdamW( 353 trainable_parameters, 354 lr=self.training_config.lr, 355 betas=(self.training_config.beta1, self.training_config.beta2), 356 eps=self.training_config.eps, 357 weight_decay=self.training_config.weight_decay, 358 fused=self.fused_adam, 359 ) 360 elif self.training_config.optimizer == "sgd": 361 self.optimizer = torch.optim.SGD( 362 trainable_parameters, 363 lr=self.training_config.lr, 364 momentum=self.training_config.momentum, 365 weight_decay=self.training_config.weight_decay, 366 nesterov=True, 367 fused=self.fused_adam, 368 ) 369 elif self.training_config.optimizer == "novograd": 370 self.optimizer = pt_optim.NovoGrad( 371 trainable_parameters, 372 lr=self.training_config.lr, 373 betas=(self.training_config.beta1, self.training_config.beta2), 374 eps=self.training_config.eps, 375 weight_decay=self.training_config.weight_decay, 376 weight_decouple=True 377 ) 378 elif self.training_config.optimizer == "fromage": 379 self.optimizer = pt_optim.Fromage( 380 trainable_parameters, 381 lr=self.training_config.lr, 382 p_bound=None, 383 ) 384 elif self.training_config.optimizer == "lamb": 385 self.optimizer = pt_optim.Lamb( 386 trainable_parameters, 387 lr=self.training_config.lr, 388 betas=(self.training_config.beta1, self.training_config.beta2), 389 eps=self.training_config.eps, 390 weight_decay=self.training_config.weight_decay, 391 weight_decouple=True, 392 rectify=True, 393 max_grad_norm=1.0, 394 ) 395 elif self.training_config.optimizer == "nero": 396 self.optimizer = pt_optim.Nero( 397 trainable_parameters, 398 lr=self.training_config.lr, 399 eps=self.training_config.eps, 400 ) 401 elif self.training_config.optimizer == "scion": 402 self.optimizer = pt_optim.SCION( 403 trainable_parameters, 404 lr=self.training_config.lr, 405 weight_decay=self.training_config.weight_decay 406 ) 407 elif self.training_config.optimizer == "apollo": 408 self.optimizer = pt_optim.APOLLO( 409 trainable_parameters, 410 lr=self.training_config.lr, 411 weight_decay=self.training_config.weight_decay 412 ) 413 else: 414 raise ValueError(f"Unsupported optimizer {self.training_config.optimizer}") 415 416 n_params = sum(p.numel() for p in self._unwrap_model().parameters()) 417 print(f"Model parameters: {n_params:,} | device: {self.device} | dtype: {self.training_config.pdtype}") 418 419 self.grad_norm_tracker = GradNormTracker(self._unwrap_model(), self.device) 420 421 # Save Checkpoint 422 423 def _checkpoint_if_needed(self) -> bool: 424 if self.trainer_state.global_step % self.training_config.save_every == 0: 425 self.checkpointer.save(self.trainer_state.global_step) 426 return True 427 return False 428 429 def save_checkpoint(self, step: int): 430 self.checkpointer.save(step) 431 432 # Evaluate 433 434 def _evaluate_if_needed(self) -> dict | None: 435 436 if self.eval_loader is None: 437 return None 438 439 if self.trainer_state.global_step % self.training_config.eval_every != 0: 440 return None 441 442 return run_eval( 443 eval_loader=self.eval_loader, 444 model=self.ctx.model, 445 training_config=self.ctx.training_config, 446 device=self.ctx.device, 447 ptdtype=self.ctx.ptdtype, 448 use_amp=self.ctx.use_amp, 449 pad_id=self.pad_id, 450 pad_token=self.pad_token, 451 eos_token=self.eos_token, 452 unk_token=self.unk_token, 453 tokenizer=self.tokenizer, 454 ) 455 456 # Set Learning Rate 457 458 def _set_lr(self, step: int) -> float: 459 lr = get_lr(step, self.training_config, self.min_lr) 460 for pg in self.optimizer.param_groups: 461 pg["lr"] = lr 462 return lr 463 464 # Training step 465 466 def training_step(self) -> Tuple[float, int, float, bool, tuple[dict,dict], dict ]: 467 468 raw_loss_acc = torch.zeros(1, device=self.device) 469 token_acc = torch.zeros(1, device=self.device) 470 471 extra_term_names = self.training_config.extra_loss_terms or [] 472 extra_term_acc = {name: torch.zeros(1, device=self.device) for name in extra_term_names} 473 474 self.optimizer.zero_grad(set_to_none=True) 475 476 # Forward and Backward Passes 477 478 for _ in range(self.training_config.grad_accum): 479 480 batch = self._get_next_batch() 481 482 input_ids = batch["input_ids"].to(self.device, non_blocking=True) 483 attention_mask = batch["attention_mask"].to(self.device, non_blocking=True) 484 485 labels = input_ids.masked_fill(attention_mask == 0, -100) 486 n_tokens = attention_mask[..., 1:].sum() 487 488 token_acc += n_tokens 489 490 with torch.autocast(device_type=self.device.type, dtype=self.ptdtype, enabled=self.use_amp): 491 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) 492 loss_scaled = outputs.loss * n_tokens 493 494 loss_scaled.backward() 495 raw_loss_acc += loss_scaled.detach() 496 497 for name in extra_term_names: 498 term_val = getattr( outputs, name, None ) 499 if term_val is not None: 500 extra_term_acc[name] += term_val 501 502 # Single CPU-GPU sync for both accumulators 503 raw_loss_sum_weighted, accum_tokens = raw_loss_acc.item(), token_acc.item() 504 step_tokens = int(accum_tokens) 505 extra_terms_sum = {name: acc.item() / self.training_config.grad_accum for name, acc in extra_term_acc.items()} 506 507 # Gradient Normalization and Stability Check 508 if accum_tokens > 0: 509 grads = [p.grad for p in self.model.parameters() if p.grad is not None] 510 torch._foreach_div_(grads, accum_tokens) 511 512 if not math.isfinite(raw_loss_sum_weighted): 513 print(f"Numerical instability at Step {self.trainer_state.global_step}. Exiting") 514 raise RuntimeError(f"Loss is {raw_loss_sum_weighted} (NaN/Inf).") 515 516 # Spike Detection and EMA Updates 517 token_weighted_step_mean_loss = raw_loss_sum_weighted / max(accum_tokens, 1) 518 is_warm = self.trainer_state.global_step >= self.training_config.spike_warmup_steps 519 520 is_spike = self.trainer_state.process_loss( 521 token_weighted_step_mean_loss, 522 self.training_config, 523 is_warm) 524 525 if is_spike and False: 526 527 self.optimizer.zero_grad(set_to_none=True) 528 if self.trainer_state.consecutive_spikes >= self.trainer_state.max_consecutive_spikes: 529 530 # self.save_checkpoint(self.trainer_state.global_step) 531 # raise RuntimeError( 532 # f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 533 # f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 534 # ) 535 536 print( 537 f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 538 f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 539 ) 540 541 print( f"[spike] loss={token_weighted_step_mean_loss:.4f} vs ema={self.trainer_state.spike_loss_ema:.4f}" ) 542 return raw_loss_sum_weighted, step_tokens, -1.0, True, ({}, {}), {} 543 544 if self.training_config.track_grad_norms: 545 agg_norms, per_layer_norms = self.grad_norm_tracker.compute() 546 else: 547 agg_norms, per_layer_norms = {}, {} 548 549 max_norm = self.training_config.grad_clip if self.training_config.grad_clip > 0.0 else float('inf') 550 grad_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm).item() 551 552 self.optimizer.step() 553 554 return raw_loss_sum_weighted, step_tokens, grad_norm, False, (agg_norms, per_layer_norms), extra_terms_sum 555 556 # Main Training loop 557 558 def train(self): 559 560 if self.trainer_state.global_step >= self.training_config.steps: 561 raise ValueError(f"global_step ({self.trainer_state.global_step}) >= steps ({self.training_config.steps}).") 562 563 self.model.train() 564 accumulator = Accumulator() 565 t0 = time.time() 566 567 for step in range(self.trainer_state.global_step, self.training_config.steps): 568 569 # Step 570 lr_now = self._set_lr(step) 571 step_outputs = self.training_step() 572 573 # Accumulate State 574 accumulator.update(*step_outputs) 575 576 self.trainer_state.global_step += 1 577 self.trainer_state.global_tokens += step_outputs[1] 578 579 # Log and Evaluate 580 if self.trainer_state.global_step % self.training_config.log_every == 0: 581 582 metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 583 eval_metrics = self._evaluate_if_needed() 584 585 log_metrics(self.ctx, self.trainer_state, metrics, eval_metrics) 586 587 accumulator = Accumulator() 588 t0 = time.time() 589 590 # Checkpoint 591 if self._checkpoint_if_needed(): 592 t0 = time.time() 593 594 # Done Training, Finalize 595 if accumulator.windowed_steps > 0: 596 597 final_metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 598 599 # Force a final eval if we have a loader 600 final_eval = run_eval( 601 eval_loader=self.eval_loader, 602 model=self.ctx.model, 603 training_config=self.ctx.training_config, 604 device=self.ctx.device, 605 ptdtype=self.ctx.ptdtype, 606 use_amp=self.ctx.use_amp, 607 pad_id=self.pad_id, 608 pad_token=self.pad_token, 609 eos_token=self.eos_token, 610 unk_token=self.unk_token, 611 tokenizer=self.tokenizer, 612 ) if self.eval_loader is not None else None 613 614 log_metrics(self.ctx, self.trainer_state, final_metrics, final_eval) 615 616 if self.aim_run is not None: 617 self.aim_run.close() 618 619 self.save_checkpoint(self.trainer_state.global_step) 620 print(f"Training complete. Final model: step_{self.trainer_state.global_step:07d}")
Trainer(config: amachine.am_transformers.am_training_config.TrainingConfig)
47 def __init__(self, config: TrainingConfig): 48 49 from aim import Run 50 51 self.training_config = config 52 53 # Environment abd Seeding 54 torch.manual_seed(self.training_config.seed) 55 if torch.cuda.is_available(): 56 torch.cuda.manual_seed_all(self.training_config.seed) 57 torch.set_float32_matmul_precision('high') 58 59 Path(self.training_config.output_dir).mkdir(exist_ok=True) 60 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") 61 62 # Save config early so we have a record even if initialization crashes 63 tr_config_path = Path(self.training_config.output_dir) / "training_config.json" 64 self.training_config.save(tr_config_path) 65 66 self.fused_adam = ( 67 self.device.type == "cuda" 68 and Version(torch.__version__.split("+")[0]) >= Version("2.0") 69 ) 70 self.ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.pdtype] 71 self.wdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[self.training_config.wdtype] 72 73 self.use_amp = self.ptdtype != torch.float32 74 self.min_lr = self.training_config.lr * self.training_config.min_lr_ratio 75 76 # Create default state counters (epoch=0, step=0, etc.) 77 self.trainer_state = TrainerState() 78 self.aim_run_hash = None 79 80 # Build the model, optimizer, and DataLoaders unconditionally. 81 self._build_components() 82 self._save_runtime_info(Path(self.training_config.output_dir)) 83 84 self.checkpointer = CheckpointManager( 85 config=self.training_config, 86 model=self.model, 87 optimizer=self.optimizer, 88 trainer_state=self.trainer_state, 89 device=self.device 90 ) 91 92 if self.training_config.resume is not None: 93 94 # This updates model, optimizer, and trainer_state in place. 95 self.checkpointer.load(self.training_config.resume) 96 self.aim_run_hash = self.trainer_state.aim_run_hash 97 98 # Fast-forward the dataset iterator if we aren't starting fresh 99 if not self.training_config.resume_reset_data: 100 self._fast_forward_dataloader(self.trainer_state.batches_consumed) 101 102 # Initialized after the checkpoint loads 103 self.aim_run = None 104 if self.training_config.aim_repo is not None: 105 106 experiment = self.training_config.experiment if self.training_config.experiment else "Train" 107 108 self.aim_run = Run( 109 repo=self.training_config.aim_repo, 110 experiment=experiment, 111 log_system_params=True, 112 run_hash=self.aim_run_hash, 113 ) 114 115 self.aim_run["hparams"] = self.training_config.to_dict() 116 if self.training_config.experiment_params is not None: 117 self.aim_run["experiment_params"] = self.training_config.experiment_params 118 119 self.ctx = TrainerContext( 120 model=self.model, 121 optimizer=self.optimizer, 122 device=self.device, 123 ptdtype=self.ptdtype, 124 use_amp=self.use_amp, 125 pad_id=self.pad_id, 126 eos_id=self.eos_id, 127 training_config=self.training_config, 128 aim_run=self.aim_run 129 )
def
training_step(self) -> Tuple[float, int, float, bool, tuple[dict, dict], dict]:
466 def training_step(self) -> Tuple[float, int, float, bool, tuple[dict,dict], dict ]: 467 468 raw_loss_acc = torch.zeros(1, device=self.device) 469 token_acc = torch.zeros(1, device=self.device) 470 471 extra_term_names = self.training_config.extra_loss_terms or [] 472 extra_term_acc = {name: torch.zeros(1, device=self.device) for name in extra_term_names} 473 474 self.optimizer.zero_grad(set_to_none=True) 475 476 # Forward and Backward Passes 477 478 for _ in range(self.training_config.grad_accum): 479 480 batch = self._get_next_batch() 481 482 input_ids = batch["input_ids"].to(self.device, non_blocking=True) 483 attention_mask = batch["attention_mask"].to(self.device, non_blocking=True) 484 485 labels = input_ids.masked_fill(attention_mask == 0, -100) 486 n_tokens = attention_mask[..., 1:].sum() 487 488 token_acc += n_tokens 489 490 with torch.autocast(device_type=self.device.type, dtype=self.ptdtype, enabled=self.use_amp): 491 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) 492 loss_scaled = outputs.loss * n_tokens 493 494 loss_scaled.backward() 495 raw_loss_acc += loss_scaled.detach() 496 497 for name in extra_term_names: 498 term_val = getattr( outputs, name, None ) 499 if term_val is not None: 500 extra_term_acc[name] += term_val 501 502 # Single CPU-GPU sync for both accumulators 503 raw_loss_sum_weighted, accum_tokens = raw_loss_acc.item(), token_acc.item() 504 step_tokens = int(accum_tokens) 505 extra_terms_sum = {name: acc.item() / self.training_config.grad_accum for name, acc in extra_term_acc.items()} 506 507 # Gradient Normalization and Stability Check 508 if accum_tokens > 0: 509 grads = [p.grad for p in self.model.parameters() if p.grad is not None] 510 torch._foreach_div_(grads, accum_tokens) 511 512 if not math.isfinite(raw_loss_sum_weighted): 513 print(f"Numerical instability at Step {self.trainer_state.global_step}. Exiting") 514 raise RuntimeError(f"Loss is {raw_loss_sum_weighted} (NaN/Inf).") 515 516 # Spike Detection and EMA Updates 517 token_weighted_step_mean_loss = raw_loss_sum_weighted / max(accum_tokens, 1) 518 is_warm = self.trainer_state.global_step >= self.training_config.spike_warmup_steps 519 520 is_spike = self.trainer_state.process_loss( 521 token_weighted_step_mean_loss, 522 self.training_config, 523 is_warm) 524 525 if is_spike and False: 526 527 self.optimizer.zero_grad(set_to_none=True) 528 if self.trainer_state.consecutive_spikes >= self.trainer_state.max_consecutive_spikes: 529 530 # self.save_checkpoint(self.trainer_state.global_step) 531 # raise RuntimeError( 532 # f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 533 # f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 534 # ) 535 536 print( 537 f"{self.trainer_state.consecutive_spikes} consecutive loss spikes detected " 538 f"(loss={token_weighted_step_mean_loss:.4f}, ema={self.trainer_state.spike_loss_ema:.4f})." 539 ) 540 541 print( f"[spike] loss={token_weighted_step_mean_loss:.4f} vs ema={self.trainer_state.spike_loss_ema:.4f}" ) 542 return raw_loss_sum_weighted, step_tokens, -1.0, True, ({}, {}), {} 543 544 if self.training_config.track_grad_norms: 545 agg_norms, per_layer_norms = self.grad_norm_tracker.compute() 546 else: 547 agg_norms, per_layer_norms = {}, {} 548 549 max_norm = self.training_config.grad_clip if self.training_config.grad_clip > 0.0 else float('inf') 550 grad_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm).item() 551 552 self.optimizer.step() 553 554 return raw_loss_sum_weighted, step_tokens, grad_norm, False, (agg_norms, per_layer_norms), extra_terms_sum
def
train(self):
558 def train(self): 559 560 if self.trainer_state.global_step >= self.training_config.steps: 561 raise ValueError(f"global_step ({self.trainer_state.global_step}) >= steps ({self.training_config.steps}).") 562 563 self.model.train() 564 accumulator = Accumulator() 565 t0 = time.time() 566 567 for step in range(self.trainer_state.global_step, self.training_config.steps): 568 569 # Step 570 lr_now = self._set_lr(step) 571 step_outputs = self.training_step() 572 573 # Accumulate State 574 accumulator.update(*step_outputs) 575 576 self.trainer_state.global_step += 1 577 self.trainer_state.global_tokens += step_outputs[1] 578 579 # Log and Evaluate 580 if self.trainer_state.global_step % self.training_config.log_every == 0: 581 582 metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 583 eval_metrics = self._evaluate_if_needed() 584 585 log_metrics(self.ctx, self.trainer_state, metrics, eval_metrics) 586 587 accumulator = Accumulator() 588 t0 = time.time() 589 590 # Checkpoint 591 if self._checkpoint_if_needed(): 592 t0 = time.time() 593 594 # Done Training, Finalize 595 if accumulator.windowed_steps > 0: 596 597 final_metrics = accumulator.to_metrics(elapsed_time=time.time() - t0, lr_now=lr_now) 598 599 # Force a final eval if we have a loader 600 final_eval = run_eval( 601 eval_loader=self.eval_loader, 602 model=self.ctx.model, 603 training_config=self.ctx.training_config, 604 device=self.ctx.device, 605 ptdtype=self.ctx.ptdtype, 606 use_amp=self.ctx.use_amp, 607 pad_id=self.pad_id, 608 pad_token=self.pad_token, 609 eos_token=self.eos_token, 610 unk_token=self.unk_token, 611 tokenizer=self.tokenizer, 612 ) if self.eval_loader is not None else None 613 614 log_metrics(self.ctx, self.trainer_state, final_metrics, final_eval) 615 616 if self.aim_run is not None: 617 self.aim_run.close() 618 619 self.save_checkpoint(self.trainer_state.global_step) 620 print(f"Training complete. Final model: step_{self.trainer_state.global_step:07d}")