GitLab Repo

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}")
 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        )
training_config
device
fused_adam
ptdtype
wdtype
use_amp
min_lr
trainer_state
aim_run_hash
checkpointer
aim_run
ctx
def save_checkpoint(self, step: int):
429    def save_checkpoint(self, step: int):
430        self.checkpointer.save(step)
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}")